/*
Copyright 2024 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.
*/
package aws
import (
"net/url"
"strings"
"github.com/gravitational/trace"
)
// IsDocumentDBEndpoint returns true if the input URI is a DocumentDB endpoint.
//
// https://docs.aws.amazon.com/documentdb/latest/developerguide/endpoints.html
func IsDocumentDBEndpoint(uri string) bool {
return isAWSServiceEndpoint(uri, DocumentDBServiceName)
}
// DocumentDBEndpointDetails contains information about a DocumentDB endpoint.
type DocumentDBEndpointDetails struct {
// ClusterID is the identifier of a DocumentDB cluster.
ClusterID string
// InstanceID is the identifier of a DocumentDB instance.
InstanceID string
// Region is the AWS region for the endpoint.
Region string
// EndpointType specifies the type of the endpoint.
EndpointType string
}
// ParseDocumentDBEndpoint parses and extracts info from the provided
// DocumentDB endpoint.
func ParseDocumentDBEndpoint(endpoint string) (*DocumentDBEndpointDetails, error) {
if !strings.HasPrefix(endpoint, "mongodb+srv://") &&
!strings.HasPrefix(endpoint, "mongodb://") {
endpoint = "mongodb://" + endpoint
}
docdbURL, err := url.Parse(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
endpoint = docdbURL.Hostname()
if strings.HasSuffix(endpoint, AWSCNEndpointSuffix) {
return parseDocumentDBCNEndpoint(endpoint)
}
if strings.HasSuffix(endpoint, AWSEndpointSuffix) {
return parseDocumentDBEndpoint(endpoint)
}
return nil, trace.BadParameter("failed to parse %v as DocumentDB endpoint", endpoint)
}
func parseDocumentDBCNEndpoint(endpoint string) (*DocumentDBEndpointDetails, error) {
// Example:
// my-documentdb-cluster-id.cluster-abcdefghijklmnop.docdb.cn-north-1.amazonaws.com.cn
parts := strings.Split(strings.TrimSuffix(endpoint, AWSCNEndpointSuffix), ".")
if len(parts) != 4 {
return nil, trace.BadParameter("failed to parse %v as DocumentDB CN endpoint", endpoint)
}
if parts[2] != DocumentDBServiceName {
return nil, trace.BadParameter("failed to parse %v as DocumentDB CN endpoint", endpoint)
}
return makeDocumentDBDetails(parts[0], parts[1], parts[3]), nil
}
func parseDocumentDBEndpoint(endpoint string) (*DocumentDBEndpointDetails, error) {
// Examples:
// my-documentdb-cluster-id.cluster-abcdefghijklmnop.us-east-1.docdb.amazonaws.com
// my-documentdb-cluster-id.cluster-ro-abcdefghijklmnop.us-east-1.docdb.amazonaws.com
// my-instance-id.abcdefghijklmnop.us-east-1.docdb.amazonaws.com
parts := strings.Split(strings.TrimSuffix(endpoint, AWSEndpointSuffix), ".")
if len(parts) != 4 {
return nil, trace.BadParameter("failed to parse %v as DocumentDB endpoint", endpoint)
}
if parts[3] != DocumentDBServiceName {
return nil, trace.BadParameter("failed to parse %v as DocumentDB endpoint", endpoint)
}
return makeDocumentDBDetails(parts[0], parts[1], parts[2]), nil
}
func makeDocumentDBDetails(id, endpointTypePart, region string) *DocumentDBEndpointDetails {
endpointType := guessDocumentDBEndpointType(endpointTypePart)
if endpointType == DocumentDBInstanceEndpoint {
return &DocumentDBEndpointDetails{
InstanceID: id,
Region: region,
EndpointType: endpointType,
}
}
return &DocumentDBEndpointDetails{
ClusterID: id,
Region: region,
EndpointType: endpointType,
}
}
func guessDocumentDBEndpointType(endpointTypePart string) string {
switch {
case strings.HasPrefix(endpointTypePart, "cluster-ro-"):
return DocumentDBClusterReaderEndpoint
case strings.HasPrefix(endpointTypePart, "cluster-"):
return DocumentDBClusterEndpoint
default:
return DocumentDBInstanceEndpoint
}
}
const (
// DocumentDBServiceName is the service name for AWS DocumentDB.
//
// TODO(greedy52) support DocumentDB Elastic clusters when IAM Auth support
// is added. Note that Elastic clusters use "docdb-elastic" as the service
// name in the endpoint.
DocumentDBServiceName = "docdb"
// DocumentDBClusterEndpoint specifies a DocumentDB primary/cluster
// endpoint.
DocumentDBClusterEndpoint = "cluster"
// DocumentDBReaderEndpoint specifies a DocumentDB reader endpoint.
DocumentDBClusterReaderEndpoint = "reader"
// DocumentDBInstanceEndpoint specifies a DocumentDB instance endpoint.
DocumentDBInstanceEndpoint = "instance"
)
/*
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.
*/
package aws
import (
"regexp"
)
// EC2 Node IDs are {AWS account ID}-{EC2 resource ID} eg:
//
// 123456789012-i-1234567890abcdef0
//
// AWS account ID is always a 12 digit number, see
//
// https://docs.aws.amazon.com/general/latest/gr/acct-identifiers.html
//
// EC2 resource ID is i-{8 or 17 hex digits}, see
//
// https://docs.aws.amazon.com/AWSEC2/latest/UserGuide/resource-ids.html
var ec2NodeIDRE = regexp.MustCompile("^[0-9]{12}-i-[0-9a-f]{8,}$")
// IsEC2NodeID returns true if the given ID looks like an EC2 node ID. Uses a
// simple regex to check. Node IDs are almost always UUIDs when set
// automatically, but can be manually overridden by admins. If someone manually
// sets a host ID that looks like one of our generated EC2 node IDs, they may be
// able to trick this function, so don't use it for any critical purpose.
func IsEC2NodeID(id string) bool {
return ec2NodeIDRE.MatchString(id)
}
/*
Copyright 2022 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.
*/
package aws
import (
"fmt"
"net"
"net/url"
"strconv"
"strings"
"github.com/gravitational/trace"
)
const maxEndpointLength = 4096
// IsAWSEndpoint returns true if the input URI is an AWS endpoint under the
// amazonaws.com domains. It deliberately excludes api.aws endpoints (see
// IsAWSAPIEndpoint) because callers like the DynamoDB/OpenSearch endpoint
// parsers only understand the amazonaws.com shapes and use this check to
// decide whether a parse failure is a config error. Once those parsers learn
// the api.aws endpoint shapes, this split may no longer be necessary.
func IsAWSEndpoint(uri string) bool {
hostname, err := removeSchemaAndPort(uri)
if err != nil {
return false
}
return strings.HasSuffix(hostname, AWSEndpointSuffix) || strings.HasSuffix(hostname, AWSCNEndpointSuffix)
}
// IsAWSAPIEndpoint returns true if the input URI is an AWS endpoint under the
// api.aws domain, used for dualstack and newer AWS service endpoints.
func IsAWSAPIEndpoint(uri string) bool {
hostname, err := removeSchemaAndPort(uri)
if err != nil {
return false
}
return strings.HasSuffix(hostname, AWSAPIEndpointSuffix)
}
// IsAWSOwnedEndpoint returns true if the input URI is under any AWS-owned
// endpoint domain (amazonaws.com, amazonaws.com.cn, or api.aws). Use this for
// trust and interception decisions. Use IsAWSEndpoint when deciding whether a
// legacy endpoint parser should have understood the URI.
//
// https://docs.aws.amazon.com/general/latest/gr/rande.html
func IsAWSOwnedEndpoint(uri string) bool {
return IsAWSEndpoint(uri) || IsAWSAPIEndpoint(uri)
}
// IsRDSEndpoint returns true if the input URI is an RDS endpoint.
//
// https://docs.aws.amazon.com/AmazonRDS/latest/AuroraUserGuide/Aurora.Overview.Endpoints.html
func IsRDSEndpoint(uri string) bool {
return isAWSServiceEndpoint(uri, RDSServiceName)
}
// IsRedshiftEndpoint returns true if the input URI is an Redshift endpoint.
//
// https://docs.aws.amazon.com/redshift/latest/mgmt/connecting-from-psql.html
func IsRedshiftEndpoint(uri string) bool {
return isAWSServiceEndpoint(uri, RedshiftServiceName)
}
// IsRedshiftServerlessEndpoint returns true if the input URI is an Redshift
// Serverless endpoint.
//
// https://docs.aws.amazon.com/redshift/latest/mgmt/serverless-connecting.html
func IsRedshiftServerlessEndpoint(uri string) bool {
return isAWSServiceEndpoint(uri, RedshiftServerlessServiceName)
}
// IsElastiCacheEndpoint returns true if the input URI is an ElastiCache
// endpoint.
func IsElastiCacheEndpoint(uri string) bool {
_, err := ParseElastiCacheEndpoint(uri)
return err == nil
}
// IsElastiCacheServerlessEndpoint returns true if the input URI is an ElastiCacheServerless
// endpoint.
func IsElastiCacheServerlessEndpoint(uri string) bool {
_, err := ParseElastiCacheServerlessEndpoint(uri)
return err == nil
}
// IsMemoryDBEndpoint returns true if the input URI is an MemoryDB
// endpoint.
func IsMemoryDBEndpoint(uri string) bool {
return isAWSServiceEndpoint(uri, MemoryDBSServiceName)
}
// IsKeyspacesEndpoint returns true if input URI is an AWS Keyspaces endpoint.
// https://docs.aws.amazon.com/keyspaces/latest/devguide/programmatic.endpoints.html
func IsKeyspacesEndpoint(uri string) bool {
hasCassandraPrefix := strings.HasPrefix(uri, "cassandra.") || strings.HasPrefix(uri, "cassandra-fips.")
return hasCassandraPrefix && IsAWSEndpoint(uri)
}
// IsOpenSearchEndpoint returns true if input URI is an OpenSearch endpoint.
func IsOpenSearchEndpoint(uri string) bool {
return isAWSServiceEndpoint(uri, OpenSearchServiceName)
}
// RDSEndpointDetails contains information about an RDS endpoint.
type RDSEndpointDetails struct {
// InstanceID is the identifier of an RDS instance.
InstanceID string
// ClusterID is the identifier of an RDS Aurora cluster.
ClusterID string
// ClusterCustomEndpointName is the identifier of an Aurora cluster custom endpoint.
ClusterCustomEndpointName string
// ProxyName is the identifier of an RDS proxy.
ProxyName string
// ProxyCustomEndpointName is the identifier of an RDS proxy custom endpoint.
ProxyCustomEndpointName string
// Region is the AWS region the database resides in.
Region string
// EndpointType specifies the type of the endpoint, if available.
//
// Note that the endpoint type of RDS Proxies are determined by their
// targets, so the endpoint type will be empty for RDS Proxies here as it
// cannot be decided by the endpoint URL itself.
EndpointType string
}
// IsProxy returns true if the RDS endpoint is an RDS Proxy.
func (d RDSEndpointDetails) IsProxy() bool {
return d.ProxyName != "" || d.ProxyCustomEndpointName != ""
}
// ParseRDSEndpoint extracts the identifier and region from the provided RDS
// endpoint.
func ParseRDSEndpoint(endpoint string) (*RDSEndpointDetails, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
if strings.ContainsRune(endpoint, ':') {
var err error
endpoint, _, err = net.SplitHostPort(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
}
if strings.HasSuffix(endpoint, AWSCNEndpointSuffix) {
return parseRDSCNEndpoint(endpoint)
}
return parseRDSEndpoint(endpoint)
}
// parseRDSEndpoint extracts the identifier and region from the provided RDS
// endpoint for standard regions.
//
// RDS/Aurora endpoints look like this:
// aurora-instance-1.abcdefghijklmnop.us-west-1.rds.amazonaws.com
func parseRDSEndpoint(endpoint string) (*RDSEndpointDetails, error) {
parts := strings.Split(endpoint, ".")
hasCorrectLen := len(parts) == 6 || len(parts) == 7
serviceNameIndex := len(parts) - 3
regionIndex := len(parts) - 4
suffixStart := regionIndex
if !strings.HasSuffix(endpoint, AWSEndpointSuffix) || !hasCorrectLen || parts[serviceNameIndex] != RDSServiceName {
return nil, trace.BadParameter("failed to parse %v as RDS endpoint", endpoint)
}
details, err := parseRDSWithoutSuffixes(endpoint, parts[:suffixStart], parts[regionIndex])
return details, trace.Wrap(err)
}
// parseRDSEndpoint extracts the identifier and region from the provided RDS
// endpoint for AWS China regions.
//
// RDS/Aurora endpoints look like this for AWS China regions:
// aurora-instance-2.abcdefghijklmnop.rds.cn-north-1.amazonaws.com.cn
func parseRDSCNEndpoint(endpoint string) (*RDSEndpointDetails, error) {
parts := strings.Split(endpoint, ".")
hasCorrectLen := len(parts) == 7 || len(parts) == 8
regionIndex := len(parts) - 4
serviceNameIndex := len(parts) - 5
suffixStart := serviceNameIndex
if !strings.HasSuffix(endpoint, AWSCNEndpointSuffix) || !hasCorrectLen || parts[serviceNameIndex] != RDSServiceName {
return nil, trace.BadParameter("failed to parse %v as RDS CN endpoint", endpoint)
}
details, err := parseRDSWithoutSuffixes(endpoint, parts[:suffixStart], parts[regionIndex])
return details, trace.Wrap(err)
}
// parseRDSWithoutSuffixes extracts identifiers from provided parts and returns
// RDSEndpointDetails. It is expected that the provided parts has either:
// - two parts (e.g. aurora-instance-1.abcdefghijklmnop)
// - or three parts (e.g. my-proxy-custom.endpoint.proxy-abcdefghijklmnop)
// as region/service/partition suffixes are removed by the caller.
func parseRDSWithoutSuffixes(endpoint string, parts []string, region string) (*RDSEndpointDetails, error) {
// RDS/Aurora instance endpoints look like this:
// aurora-instance-1.abcdefghijklmnop.<suffixes>
//
// Aurora cluster endpoints look like this:
// my-cluster.cluster-abcdefghijklmnop.<suffixes>
// my-cluster.cluster-ro-abcdefghijklmnop.<suffixes>
// my-custom.cluster-custom-abcdefghijklmnop.<suffixes>
//
// RDS Proxy "default" endpoints look like this:
// my-proxy.proxy-abcdefghijklmnop.<suffixes>
//
// RDS Proxy custom endpoints look like this:
// my-proxy-custom.endpoint.proxy-abcdefghijklmnop.<suffixes>
//
// https://docs.aws.amazon.com/AmazonRDS/latest/AuroraUserGuide/Aurora.Overview.Endpoints.html
// https://docs.aws.amazon.com/AmazonRDS/latest/UserGuide/rds-proxy-setup.html#rds-proxy-connecting
// https://docs.aws.amazon.com/AmazonRDS/latest/UserGuide/rds-proxy-endpoints.html
switch len(parts) {
case 2:
switch {
case strings.HasPrefix(parts[1], "cluster-custom-"):
// Note that we are not able to get the cluster ID from the cluster
// custom endpoints. The cluster ID must be provided separately in
// addition to the endpoints.
return &RDSEndpointDetails{
ClusterCustomEndpointName: parts[0],
Region: region,
EndpointType: RDSEndpointTypeCustom,
}, nil
case strings.HasPrefix(parts[1], "cluster-ro-"):
return &RDSEndpointDetails{
ClusterID: parts[0],
Region: region,
EndpointType: RDSEndpointTypeReader,
}, nil
case strings.HasPrefix(parts[1], "cluster-"):
return &RDSEndpointDetails{
ClusterID: parts[0],
Region: region,
EndpointType: RDSEndpointTypePrimary,
}, nil
case strings.HasPrefix(parts[1], "proxy-"):
return &RDSEndpointDetails{
ProxyName: parts[0],
Region: region,
}, nil
default:
return &RDSEndpointDetails{
InstanceID: parts[0],
Region: region,
EndpointType: RDSEndpointTypeInstance,
}, nil
}
case 3:
if strings.HasPrefix(parts[2], "proxy-") && parts[1] == "endpoint" {
return &RDSEndpointDetails{
ProxyCustomEndpointName: parts[0],
Region: region,
}, nil
}
return nil, trace.BadParameter("failed to parse %v as RDS Proxy custom endpoint", endpoint)
default:
return nil, trace.BadParameter("failed to parse %v as RDS endpoint", endpoint)
}
}
// ParseRedshiftEndpoint extracts cluster ID and region from the provided
// Redshift endpoint.
func ParseRedshiftEndpoint(endpoint string) (clusterID, region string, err error) {
if len(endpoint) > maxEndpointLength {
return "", "", trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
if strings.ContainsRune(endpoint, ':') {
endpoint, _, err = net.SplitHostPort(endpoint)
if err != nil {
return "", "", trace.Wrap(err)
}
}
if strings.HasSuffix(endpoint, AWSCNEndpointSuffix) {
return parseRedshiftCNEndpoint(endpoint)
}
return parseRedshiftEndpoint(endpoint)
}
// parseRedshiftEndpoint extracts cluster ID and region from the provided
// Redshift endpoint for standard regions.
//
// Redshift endpoints look like this:
// redshift-cluster-1.abcdefghijklmnop.us-east-1.redshift.amazonaws.com
func parseRedshiftEndpoint(endpoint string) (clusterID, region string, err error) {
parts := strings.Split(endpoint, ".")
if !strings.HasSuffix(endpoint, AWSEndpointSuffix) || len(parts) != 6 || parts[3] != RedshiftServiceName {
return "", "", trace.BadParameter("failed to parse %v as Redshift endpoint", endpoint)
}
return parts[0], parts[2], nil
}
// parseRedshiftCNEndpoint extracts cluster ID and region from the provided
// Redshift endpoint for AWS China regions.
//
// Redshift endpoints look like this for AWS China regions:
// redshift-cluster-2.abcdefghijklmnop.redshift.cn-north-1.amazonaws.com.cn
func parseRedshiftCNEndpoint(endpoint string) (clusterID, region string, err error) {
parts := strings.Split(endpoint, ".")
if !strings.HasSuffix(endpoint, AWSCNEndpointSuffix) || len(parts) != 7 || parts[2] != RedshiftServiceName {
return "", "", trace.BadParameter("failed to parse %v as Redshift CN endpoint", endpoint)
}
return parts[0], parts[3], nil
}
// RedshiftServerlessEndpointDetails contains information about an Redshift
// Serverless endpoint.
type RedshiftServerlessEndpointDetails struct {
// WorkgroupName is the name of the workgroup.
WorkgroupName string
// EndpointName is the name of the VPC endpoint.
EndpointName string
// AccountID is the AWS Account ID.
AccountID string
// Region is the AWS region the database resides in.
Region string
}
// ParseRedshiftServerlessEndpoint extracts name, AWS Account ID, and region
// from the provided Redshift Serverless endpoint.
func ParseRedshiftServerlessEndpoint(endpoint string) (*RedshiftServerlessEndpointDetails, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
if strings.ContainsRune(endpoint, ':') {
var err error
endpoint, _, err = net.SplitHostPort(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
}
if strings.HasSuffix(endpoint, AWSCNEndpointSuffix) {
// TODO(greedy52) add AWS China support when Redshift Serverless come to those regions.
return nil, trace.NotImplemented("failed to parse %v as Redshift Serverless endpoint: AWS China regions are not supported yet", endpoint)
}
return parseRedshiftServerlessEndpoint(endpoint)
}
// parseRedshiftServerlessEndpoint extracts name, AWS account ID, and region
// from the provided Redshift Serverless endpoint for standard regions.
//
// Workgroup endpoint looks like this:
// <workgroup-name>.<account-id>.<region>.redshift-serverless.amazonaws.com
//
// VPC endpoint looks like this:
// <vpc-endpoint-name>-endpoint-<some-hash>.<account-id>.<region>.redshift-serverless.amazonaws.com
func parseRedshiftServerlessEndpoint(endpoint string) (*RedshiftServerlessEndpointDetails, error) {
parts := strings.Split(endpoint, ".")
if !strings.HasSuffix(endpoint, AWSEndpointSuffix) || len(parts) != 6 || parts[3] != RedshiftServerlessServiceName {
return nil, trace.BadParameter("failed to parse %v as Redshift Serverless endpoint", endpoint)
}
if endpointName, _, found := strings.Cut(parts[0], "-endpoint-"); found {
return &RedshiftServerlessEndpointDetails{
EndpointName: endpointName,
AccountID: parts[1],
Region: parts[2],
}, nil
}
return &RedshiftServerlessEndpointDetails{
WorkgroupName: parts[0],
AccountID: parts[1],
Region: parts[2],
}, nil
}
// RedisEndpointInfo describes details extracted from a ElastiCache or MemoryDB
// Redis endpoint.
type RedisEndpointInfo struct {
// ID is the identifier of the endpoint.
ID string
// Region is the AWS region for the endpoint.
Region string
// TransitEncryptionEnabled specifies if in-transit encryption (TLS) is
// enabled.
TransitEncryptionEnabled bool
// EndpointType specifies the type of the endpoint.
EndpointType string
}
const (
// ElastiCacheConfigurationEndpoint is the configuration endpoint that used
// for cluster mode connection.
ElastiCacheConfigurationEndpoint = "configuration"
// ElastiCachePrimaryEndpoint is the endpoint of the primary node in the
// node group.
ElastiCachePrimaryEndpoint = "primary"
// ElastiCacheReaderEndpoint is the endpoint of the replica nodes in the
// node group.
ElastiCacheReaderEndpoint = "reader"
// ElastiCacheNodeEndpoint is the endpoint that used to connect to an
// individual node.
ElastiCacheNodeEndpoint = "node"
// MemoryDBClusterEndpoint is the cluster configuration endpoint for a
// MemoryDB cluster.
MemoryDBClusterEndpoint = "cluster"
// MemoryDBNodeEndpoint is the endpoint of an individual MemoryDB node.
MemoryDBNodeEndpoint = "node"
// OpenSearchDefaultEndpoint is the default endpoint for domain.
OpenSearchDefaultEndpoint = "default"
// OpenSearchCustomEndpoint is the custom endpoint configured for domain.
OpenSearchCustomEndpoint = "custom"
// OpenSearchVPCEndpoint is the VPC endpoint for domain.
OpenSearchVPCEndpoint = "vpc"
// RDSEndpointTypePrimary is the endpoint that specifies the connection for
// the primary instance of the RDS cluster.
RDSEndpointTypePrimary = "primary"
// RDSEndpointTypeReader is the endpoint that load-balances connections
// across the Aurora Replicas that are available in an RDS cluster.
RDSEndpointTypeReader = "reader"
// RDSEndpointTypeCustom is the endpoint that specifies one of the custom
// endpoints associated with the RDS cluster.
RDSEndpointTypeCustom = "custom"
// RDSEndpointTypeInstance is the endpoint of an RDS DB instance.
RDSEndpointTypeInstance = "instance"
)
// ParseElastiCacheEndpoint extracts the details from the provided
// ElastiCache Redis endpoint.
//
// https://docs.aws.amazon.com/AmazonElastiCache/latest/red-ug/GettingStarted.ConnectToCacheNode.html
func ParseElastiCacheEndpoint(endpoint string) (*RedisEndpointInfo, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
endpoint, err := removeSchemaAndPort(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
// Remove partition suffix. Note that endpoints for CN regions use the same
// format except they end with AWSCNEndpointSuffix.
endpointWithoutSuffix, _, err := removePartitionSuffix(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
// Split into parts to extract details. They look like this in general:
// <part>.<part>.<part>.<short-region>.cache
//
// Note that ElastiCache uses short region codes like "use1".
//
// For Redis with cluster mode enabled, users can connect through either
// "configuration" endpoint or individual "node" endpoints.
// For Redis with cluster mode disabled, users can connect through either
// "primary", "reader", or individual "node" endpoints.
parts := strings.Split(endpointWithoutSuffix, ".")
if len(parts) == 5 && parts[4] == ElastiCacheServiceName {
region, ok := ShortRegionToRegion(parts[3])
if !ok {
return nil, trace.BadParameter("%v is not a valid region", parts[3])
}
// Configuration endpoint for Redis with TLS enabled looks like:
// clustercfg.my-redis-shards.xxxxxx.use1.cache.<suffix>:6379
if parts[0] == "clustercfg" {
return &RedisEndpointInfo{
ID: parts[1],
Region: region,
TransitEncryptionEnabled: true,
EndpointType: ElastiCacheConfigurationEndpoint,
}, nil
}
// Configuration endpoint for Redis with TLS disabled looks like:
// my-redis-shards.xxxxxx.clustercfg.use1.cache.<suffix>:6379
if parts[2] == "clustercfg" {
return &RedisEndpointInfo{
ID: parts[0],
Region: region,
TransitEncryptionEnabled: false,
EndpointType: ElastiCacheConfigurationEndpoint,
}, nil
}
// Node endpoint for Redis with TLS disabled looks like:
// my-redis-cluster-001.xxxxxx.0001.use0.cache.<suffix>:6379
// my-redis-shards-0001-001.xxxxxx.0001.use0.cache.<suffix>:6379
if isElasticCacheShardID(parts[2]) {
return &RedisEndpointInfo{
ID: trimElastiCacheShardAndNodeID(parts[0]),
Region: region,
TransitEncryptionEnabled: false,
EndpointType: ElastiCacheNodeEndpoint,
}, nil
}
// Node, primary, reader endpoints for Redis with TLS enabled look like:
// my-redis-cluster-001.my-redis-cluster.xxxxxx.use1.cache.<suffix>:6379
// my-redis-shards-0001-001.my-redis-shards.xxxxxx.use1.cache.<suffix>:6379
// master.my-redis-cluster.xxxxxx.use1.cache.<suffix>:6379
// replica.my-redis-cluster.xxxxxx.use1.cache.<suffix>:6379
var endpointType string
switch strings.ToLower(parts[0]) {
case "master":
endpointType = ElastiCachePrimaryEndpoint
case "replica":
endpointType = ElastiCacheReaderEndpoint
default:
endpointType = ElastiCacheNodeEndpoint
}
return &RedisEndpointInfo{
ID: parts[1],
Region: region,
TransitEncryptionEnabled: true,
EndpointType: endpointType,
}, nil
}
// Primary and reader endpoints for Redis with TLS disabled have an extra
// shard ID in the endpoints, and they look like:
// my-redis-cluster.xxxxxx.ng.0001.use1.cache.<suffix>:6379
// my-redis-cluster-ro.xxxxxx.ng.0001.use1.cache.<suffix>:6379
if len(parts) == 6 && parts[5] == ElastiCacheServiceName && isElasticCacheShardID(parts[3]) {
region, ok := ShortRegionToRegion(parts[4])
if !ok {
return nil, trace.BadParameter("%v is not a valid region", parts[4])
}
// Remove "-ro" from reader endpoint.
if before, ok := strings.CutSuffix(parts[0], "-ro"); ok {
return &RedisEndpointInfo{
ID: before,
Region: region,
TransitEncryptionEnabled: false,
EndpointType: ElastiCacheReaderEndpoint,
}, nil
}
return &RedisEndpointInfo{
ID: parts[0],
Region: region,
TransitEncryptionEnabled: false,
EndpointType: ElastiCachePrimaryEndpoint,
}, nil
}
return nil, trace.BadParameter("unknown ElastiCache Redis endpoint format %q", endpoint)
}
// isElasticCacheShardID returns true if the input part is in shard ID format.
// The shard ID is a 4 character string of an integer left padded with zeros
// (e.g. 0001).
func isElasticCacheShardID(part string) bool {
if len(part) != 4 {
return false
}
_, err := strconv.Atoi(part)
return err == nil
}
// isElasticCacheNodeID returns true if the input part is in node ID format.
// The node ID is a 3 character string of an integer left padded with zeros
// (e.g. 001).
func isElasticCacheNodeID(part string) bool {
if len(part) != 3 {
return false
}
_, err := strconv.Atoi(part)
return err == nil
}
// trimElastiCacheShardAndNodeID trims shard and node ID suffix from input.
func trimElastiCacheShardAndNodeID(input string) string {
// input can be one of:
// <replication-group-id>
// <replication-group-id>-<node-id>
// <replication-group-id>-<shard-id>-<node-id>
parts := strings.Split(input, "-")
if len(parts) > 0 {
if isElasticCacheNodeID(parts[len(parts)-1]) {
parts = parts[:len(parts)-1]
}
}
if len(parts) > 0 {
if isElasticCacheShardID(parts[len(parts)-1]) {
parts = parts[:len(parts)-1]
}
}
return strings.Join(parts, "-")
}
// ParseElastiCacheServerlessEndpoint extracts the details from the provided
// ElastiCacheServerless Redis endpoint, which should be in the form
// <cache_name>.serverless.<region>.cache.amazonaws.com:<port>
func ParseElastiCacheServerlessEndpoint(endpoint string) (*RedisEndpointInfo, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
endpoint, err := removeSchemaAndPort(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
// Remove partition suffix. Note that endpoints for CN regions use the same
// format except they end with AWSCNEndpointSuffix.
endpointWithoutSuffix, _, err := removePartitionSuffix(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
// Split into parts to extract details. They look like this in general:
// <cache_name>.serverless.<short-region>.cache
//
// Note that ElastiCache uses short region codes like "use1".
const (
nameIdx = iota
serverlessIdx
shortRegionIdx
svcIdx
numParts
)
parts := strings.Split(endpointWithoutSuffix, ".")
if len(parts) != numParts || parts[serverlessIdx] != "serverless" || parts[svcIdx] != ElastiCacheServiceName {
return nil, trace.BadParameter("unknown ElastiCache Redis endpoint format %q", endpoint)
}
region, ok := ShortRegionToRegion(parts[shortRegionIdx])
if !ok {
return nil, trace.BadParameter("%v is not a valid region", parts[shortRegionIdx])
}
// Configuration endpoint for Redis with TLS disabled looks like:
// example-<randomhex>.serverless.cac1.cache.amazonaws.com:6379
info := &RedisEndpointInfo{
ID: parts[nameIdx],
Region: region,
TransitEncryptionEnabled: true,
EndpointType: ElastiCacheConfigurationEndpoint,
}
nameParts := strings.Split(info.ID, "-")
if len(nameParts) > 1 {
info.ID = strings.Join(nameParts[:len(nameParts)-1], "-")
}
return info, nil
}
// ParseMemoryDBEndpoint extracts the details from the provided
// MemoryDB endpoint.
//
// https://docs.aws.amazon.com/memorydb/latest/devguide/endpoints.html
func ParseMemoryDBEndpoint(endpoint string) (*RedisEndpointInfo, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
endpoint, err := removeSchemaAndPort(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
// Here is a sample endpoint for MemoryDB:
// clustercfg.my-memorydb.scwzlu.memorydb.ca-central-1.amazonaws.com
//
// Unlike RDS/Redshift endpoints, the service subdomain is before region.
// Unlike ElastiCache endpoints, MemoryDB uses full region name.
endpointWithoutSuffix, _, err := removePartitionSuffix(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
parts := strings.Split(endpointWithoutSuffix, ".")
if len(parts) != 5 || parts[3] != MemoryDBSServiceName {
return nil, trace.BadParameter("unknown MemoryDB endpoint format")
}
switch {
// TLS disabled cluster endpoints look like this:
// <cluster-name>.<xxxx>.clustercfg.memorydb.<region>.<suffix>
case parts[2] == "clustercfg":
return &RedisEndpointInfo{
ID: parts[0],
Region: parts[4],
TransitEncryptionEnabled: false,
EndpointType: MemoryDBClusterEndpoint,
}, nil
// TLS enabled cluster endpoints look like this:
// clustercfg.<cluster-name>.<xxxx>.memorydb.<region>.<suffix>
case parts[0] == "clustercfg":
return &RedisEndpointInfo{
ID: parts[1],
Region: parts[4],
TransitEncryptionEnabled: true,
EndpointType: MemoryDBClusterEndpoint,
}, nil
// TLS disabled node endpoints look like this:
// <cluster-name>-<shard-id>-<node-id>.<xxxx>.<shard-id>.memorydb.<region>.<suffix>
//
// MemoryDB and ElastiCache share same shard/node ID format.
case isElasticCacheShardID(parts[2]):
return &RedisEndpointInfo{
ID: trimElastiCacheShardAndNodeID(parts[0]),
Region: parts[4],
TransitEncryptionEnabled: false,
EndpointType: MemoryDBNodeEndpoint,
}, nil
// TLS enabled node endpoints look like this:
// <cluster-name>-<shard-id>-<node-id>.<cluster-name>.<xxxx>.memorydb.<region>.<suffix>
default:
return &RedisEndpointInfo{
ID: trimElastiCacheShardAndNodeID(parts[0]),
Region: parts[4],
TransitEncryptionEnabled: true,
EndpointType: MemoryDBNodeEndpoint,
}, nil
}
}
// isAWSServiceEndpoint returns true if uri is a valid AWS endpoint and uri
// contains the provided service name as a subdomain.
func isAWSServiceEndpoint(uri, serviceName string) bool {
return IsAWSEndpoint(uri) && strings.Contains(uri, fmt.Sprintf(".%s.", serviceName))
}
func removeSchemaAndPort(endpoint string) (string, error) {
// Add a temporary schema to make a valid URL for url.Parse.
if !strings.Contains(endpoint, "://") {
endpoint = "schema://" + endpoint
}
parsedURL, err := url.Parse(endpoint)
if err != nil {
return "", trace.Wrap(err)
}
return parsedURL.Hostname(), nil
}
func removePartitionSuffix(endpoint string) (string, string, error) {
switch {
case strings.HasSuffix(endpoint, AWSEndpointSuffix):
return strings.TrimSuffix(endpoint, AWSEndpointSuffix), AWSEndpointSuffix, nil
case strings.HasSuffix(endpoint, AWSCNEndpointSuffix):
return strings.TrimSuffix(endpoint, AWSCNEndpointSuffix), AWSCNEndpointSuffix, nil
default:
return "", "", trace.BadParameter("%v is not a valid AWS endpoint", endpoint)
}
}
const (
// AWSEndpointSuffix is the endpoint suffix for AWS Standard and AWS US
// GovCloud regions.
//
// https://docs.aws.amazon.com/general/latest/gr/rande.html#regional-endpoints
// https://docs.aws.amazon.com/govcloud-us/latest/UserGuide/using-govcloud-endpoints.html
AWSEndpointSuffix = ".amazonaws.com"
// AWSCNEndpointSuffix is the endpoint suffix for AWS China regions.
//
// https://docs.amazonaws.cn/en_us/aws/latest/userguide/endpoints-arns.html
AWSCNEndpointSuffix = ".amazonaws.com.cn"
// AWSAPIEndpointSuffix is the endpoint suffix for AWS dualstack and newer
// service endpoints. The api.aws domain is owned and operated by AWS.
//
// https://docs.aws.amazon.com/general/latest/gr/rande.html#dual-stack-endpoints
AWSAPIEndpointSuffix = ".api.aws"
// RDSServiceName is the service name for AWS RDS.
RDSServiceName = "rds"
// RedshiftServiceName is the service name for AWS Redshift.
RedshiftServiceName = "redshift"
// RedshiftServerlessServiceName is the service name for AWS Redshift Serverless.
RedshiftServerlessServiceName = "redshift-serverless"
// ElastiCacheServiceName is the service name for AWS ElastiCache.
ElastiCacheServiceName = "cache"
// MemoryDBSServiceName is the service name for AWS MemoryDB.
MemoryDBSServiceName = "memorydb"
// DynamoDBServiceName is the service name for AWS DynamoDB.
DynamoDBServiceName = "dynamodb"
// DynamoDBFipsServiceName is the fips variant service name for AWS DynamoDB.
DynamoDBFipsServiceName = "dynamodb-fips"
// DynamoDBStreamsServiceName is the AWS DynamoDB Streams service name.
DynamoDBStreamsServiceName = "streams.dynamodb"
// DAXServiceName is the AWS DynamoDB Accelerator service name.
DAXServiceName = "dax"
// OpenSearchServiceName is the AWS OpenSearch service name.
OpenSearchServiceName = "es"
)
// CassandraEndpointURLForRegion returns a Cassandra endpoint based on the provided region.
// https://docs.aws.amazon.com/keyspaces/latest/devguide/programmatic.endpoints.html
func CassandraEndpointURLForRegion(region string) string {
if IsCNRegion(region) {
return fmt.Sprintf("cassandra.%s%s:9142", region, AWSCNEndpointSuffix)
}
return fmt.Sprintf("cassandra.%s%s:9142", region, AWSEndpointSuffix)
}
// CassandraEndpointRegion returns an AWS region from cassandra endpoint:
// where endpoint looks like cassandra.us-east-2.amazonaws.com
// https://docs.aws.amazon.com/keyspaces/latest/devguide/programmatic.endpoints.html
func CassandraEndpointRegion(endpoint string) (string, error) {
parts, _, err := extractAWSEndpointParts(endpoint)
if err != nil {
return "", trace.Wrap(err)
}
if len(parts) != 2 {
return "", trace.BadParameter("invalid Cassandra endpoint")
}
return parts[1], nil
}
// DynamoDBEndpointInfo describes info extracted from a DynamoDB endpoint.
type DynamoDBEndpointInfo struct {
// Service is the service subdomain of the endpoint, for example "dynamodb" or "dax".
Service string
// Region is the AWS region for the endpoint, for example "us-west-1".
Region string
// Partition is the AWS partition for the endpoint, for example ".amazonaws.com"
Partition string
}
// ParseDynamoDBEndpoint parses and extract info from the provided DynamoDB endpoint.
func ParseDynamoDBEndpoint(endpoint string) (*DynamoDBEndpointInfo, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
endpoint = strings.ToLower(endpoint)
parts, partition, err := extractAWSEndpointParts(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
switch len(parts) {
case 2, 3:
default:
return nil, trace.BadParameter("invalid DynamoDB endpoint %q", endpoint)
}
info := &DynamoDBEndpointInfo{
Service: strings.Join(parts[:len(parts)-1], "."),
Region: parts[len(parts)-1],
Partition: partition,
}
// check for recognized service name.
switch info.Service {
case DynamoDBServiceName, DynamoDBFipsServiceName,
DynamoDBStreamsServiceName, DAXServiceName:
default:
return nil, trace.BadParameter("invalid DynamoDB endpoint %q", endpoint)
}
// check that the partition is valid for the region.
if info.Region == "" || info.Partition == "" {
return nil, trace.BadParameter("invalid DynamoDB endpoint %q", endpoint)
}
switch {
case info.Partition == AWSCNEndpointSuffix && IsCNRegion(info.Region):
case info.Partition == AWSEndpointSuffix && !IsCNRegion(info.Region):
default:
return nil, trace.BadParameter("invalid AWS region %q for AWS partition %q",
info.Region, info.Partition)
}
return info, nil
}
// OpenSearchEndpointInfo describes info extracted from an AWS endpoint.
type OpenSearchEndpointInfo struct {
// Service is the service subdomain of the endpoint. Only "es" allowed for now.
Service string
// Region is the AWS region for the endpoint, for example "us-west-1".
Region string
// Partition is the AWS partition for the endpoint, for example ".amazonaws.com"
Partition string
}
// ParseOpensearchEndpoint parses and extract info from the provided OpenSearch endpoint.
func ParseOpensearchEndpoint(endpoint string) (*OpenSearchEndpointInfo, error) {
if len(endpoint) > maxEndpointLength {
return nil, trace.BadParameter("invalid endpoint exceeds maximum length of %d", maxEndpointLength)
}
endpoint = strings.ToLower(endpoint)
parts, partition, err := extractAWSEndpointParts(endpoint)
if err != nil {
return nil, trace.Wrap(err)
}
if len(parts) != 3 {
return nil, trace.BadParameter("invalid OpenSearch endpoint %q, wrong number of parts %v", endpoint, len(parts))
}
info := &OpenSearchEndpointInfo{
Region: parts[len(parts)-2],
Service: parts[len(parts)-1],
Partition: partition,
}
// check for recognized service name.
if info.Service != OpenSearchServiceName {
return nil, trace.BadParameter("invalid OpenSearch endpoint %q, invalid service %q", endpoint, info.Service)
}
// check that the partition is valid for the region.
switch {
case info.Region == "" || info.Partition == "":
return nil, trace.BadParameter("invalid OpenSearch endpoint %q, empty partition and region", endpoint)
case info.Region == "":
return nil, trace.BadParameter("invalid OpenSearch endpoint %q, empty region", endpoint)
case info.Partition == "":
return nil, trace.BadParameter("invalid OpenSearch endpoint %q, empty partition", endpoint)
}
switch {
case info.Partition == AWSCNEndpointSuffix && IsCNRegion(info.Region):
case info.Partition == AWSEndpointSuffix && !IsCNRegion(info.Region):
default:
return nil, trace.BadParameter("invalid AWS region %q for AWS partition %q",
info.Region, info.Partition)
}
return info, nil
}
// DynamoDBURIForRegion constructs a DynamoDB URI based on the AWS region.
// The URI uses a custom schema aws:// to differentiate an auto-generated URI from
// a user-configured URI in the engine.
// When the Teleport DynamoDB engine sees this custom URI schema, it will resolve
// the real endpoint using the request API target.
// https://docs.aws.amazon.com/general/latest/gr/ddb.html
func DynamoDBURIForRegion(region string) string {
var suffix string
if IsCNRegion(region) {
suffix = AWSCNEndpointSuffix
} else {
suffix = AWSEndpointSuffix
}
return fmt.Sprintf("aws://dynamodb.%s%s", region, suffix)
}
// extractAWSEndpointParts strips the schema, port, and AWS suffix,
// then splits the prefix by subdomain separator (".") and returns the parts and suffix.
func extractAWSEndpointParts(endpoint string) ([]string, string, error) {
uri, err := removeSchemaAndPort(endpoint)
if err != nil {
return nil, "", trace.Wrap(err)
}
prefix, suffix, err := removePartitionSuffix(uri)
if err != nil {
return nil, "", trace.Wrap(err)
}
return strings.Split(prefix, "."), suffix, nil
}
/*
Copyright 2022 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.
*/
package aws
import (
"regexp"
"slices"
"strings"
"github.com/gravitational/trace"
)
// IsValidAccountID checks whether the accountID is a valid AWS Account ID
//
// https://docs.aws.amazon.com/accounts/latest/reference/manage-acct-identifiers.html
func IsValidAccountID(accountID string) error {
if len(accountID) != 12 {
return trace.BadParameter("must be 12-digit")
}
for _, d := range accountID {
if d < '0' || d > '9' {
return trace.BadParameter("must be 12-digit")
}
}
return nil
}
// IsValidIAMRoleName checks whether the role name is a valid AWS IAM Role identifier.
//
// > Length Constraints: Minimum length of 1. Maximum length of 64.
// > Pattern: [\w+=,.@-]+
// https://docs.aws.amazon.com/IAM/latest/APIReference/API_CreateRole.html
func IsValidIAMRoleName(roleName string) error {
if len(roleName) == 0 || len(roleName) > 64 || !matchRoleName(roleName) {
return trace.BadParameter("role is invalid")
}
return nil
}
// IsValidIAMRolesAnywhereTrustAnchorName checks whether the AWS IAM Roles Anywhere Trust Anchor name is valid.
// Validation based on the AWS documentation.
// See https://docs.aws.amazon.com/rolesanywhere/latest/APIReference/API_CreateTrustAnchor.html#API_CreateTrustAnchor_RequestBody
func IsValidIAMRolesAnywhereTrustAnchorName(name string) error {
if !matchRolesAnywhereTrustAnchorName(name) {
return trace.BadParameter("trust anchor name is invalid")
}
return nil
}
// IsValidIAMRolesAnywhereProfileName checks whether the AWS IAM Roles Anywhere Profile name is valid.
// Validation based on the AWS documentation.
// See https://docs.aws.amazon.com/rolesanywhere/latest/APIReference/API_CreateProfile.html#API_CreateProfile_RequestBody
func IsValidIAMRolesAnywhereProfileName(name string) error {
if !matchRolesAnywhereProfileName(name) {
return trace.BadParameter("profile name is invalid")
}
return nil
}
// IsValidIAMPolicyName checks whether the policy name is a valid AWS IAM Policy
// identifier.
//
// > Length Constraints: Minimum length of 1. Maximum length of 128.
// > Pattern: [\w+=,.@-]+
// https://docs.aws.amazon.com/IAM/latest/APIReference/API_CreatePolicy.html
func IsValidIAMPolicyName(policyName string) error {
// The same regex is used for role and policy names.
if len(policyName) == 0 || len(policyName) > 128 || !matchRoleName(policyName) {
return trace.BadParameter("policy name is invalid")
}
return nil
}
const (
// AWSGlobalRegion is a sentinel value used by AWS to be able to use global endpoints, instead of region specific ones.
// Useful for STS API Calls.
// https://docs.aws.amazon.com/sdkref/latest/guide/feature-region.html
AWSGlobalRegion = "aws-global"
)
// IsValidRegion ensures the region looks to be valid.
// It does not do a full validation, because AWS doesn't provide documentation for that.
// However, they usually only have the following chars: [a-z0-9\-]
func IsValidRegion(region string) error {
if region == AWSGlobalRegion {
return nil
}
if matchRegion.MatchString(region) {
return nil
}
return trace.BadParameter("region %q is invalid", region)
}
// IsValidRegionWithWeakCheck validates a region accepting any string of
// lowercase letters, digits, and hyphens. Use this instead of IsValidRegion
// when customers can be interacting with their own AWS like API implementation
// like S3 where regions do not follow strict naming convention.
func IsValidRegionWithWeakCheck(region string) error {
if weakRegionValidation.MatchString(region) {
return nil
}
return trace.BadParameter("region %q is invalid", region)
}
// IsValidPartition checks if partition is a valid AWS partition
func IsValidPartition(partition string) error {
if slices.Contains(validPartitions, partition) {
return nil
}
return trace.BadParameter("partition %q is invalid", partition)
}
// IsValidAthenaWorkgroupName checks whether the name is a valid AWS Athena
// workgroup name.
func IsValidAthenaWorkgroupName(workgroup string) error {
if matchAthenaWorkgroupName(workgroup) {
return nil
}
return trace.BadParameter("athena workgroup name %q is invalid", workgroup)
}
// IsValidGlueResourceName check whether the name is valid for an AWS Glue
// database or table used with AWS Athena
func IsValidGlueResourceName(name string) error {
if matchGlueName(name) {
return nil
}
return trace.BadParameter("glue resource name %q is invalid", name)
}
const (
arnDelimiter = ":"
arnPrefix = "arn:"
arnSections = 6
sectionService = 2 // arn:<partition>:<service>:...
sectionAccount = 4 // arn:<partition>:<service>:<region>:<accountid>:...
sectionResource = 5 // arn:<partition>:<service>:<region>:<accountid>:<resource>
iamServiceName = "iam"
)
// CheckRoleARN returns whether a string is a valid IAM Role ARN.
// Example role ARN: arn:aws:iam::123456789012:role/some-role-name
func CheckRoleARN(arn string) error {
if !strings.HasPrefix(arn, arnPrefix) {
return trace.BadParameter("arn: invalid prefix: %q", arn)
}
sections := strings.SplitN(arn, arnDelimiter, arnSections)
if len(sections) != arnSections {
return trace.BadParameter("arn: not enough sections: %q", arn)
}
resourceParts := strings.SplitN(sections[sectionResource], "/", 2)
if resourceParts[0] != "role" || sections[sectionService] != iamServiceName {
return trace.BadParameter("%q is not an AWS IAM role ARN", arn)
}
if len(resourceParts) < 2 || resourceParts[1] == "" {
return trace.BadParameter("%q is missing AWS IAM role name", arn)
}
if err := IsValidAccountID(sections[sectionAccount]); err != nil {
return trace.BadParameter("%q invalid account ID: %v", arn, err)
}
return nil
}
var (
// matchRoleName is a regex that matches against AWS IAM Role Names.
matchRoleName = regexp.MustCompile(`^[\w+=,.@-]+$`).MatchString
// matchRegion is a regex that defines the format of AWS regions.
//
// The regex matches the following from left to right:
// - starts with 2 lower case letters that represents a geo region like a
// country code
// - optional -gov, -iso, -isob for corresponding partitions
// - a word that should be a direction like "east", "west", etc.
// - a number counter
//
// Reference:
// https://github.com/aws/aws-sdk-go-v2/blob/main/codegen/smithy-aws-go-codegen/src/main/resources/software/amazon/smithy/aws/go/codegen/endpoints.json
matchRegion = regexp.MustCompile(`^(eusc-)?[a-z]{2}(-gov|-iso|-isob|-isoe|-isof)?-\w+-\d+$`)
// weakRegionValidation is a permissive regex used as a fallback when a
// region does not match the stricter matchRegion pattern. It accepts any
// string composed solely of lowercase letters, digits, and hyphens, which
// covers regions that may not yet conform to the known naming
// convention (e.g. users that have AWS compatible services). This prevents
// invalid input (special characters, whitespace, injection
// attempts).
weakRegionValidation = regexp.MustCompile(`^[a-zA-Z0-9-_]+$`)
// https://docs.aws.amazon.com/athena/latest/APIReference/API_CreateWorkGroup.html
matchAthenaWorkgroupName = regexp.MustCompile(`^[a-zA-Z0-9._-]{1,128}$`).MatchString
// https://docs.aws.amazon.com/athena/latest/ug/tables-databases-columns-names.html
// More strict than strictly necessary, but a good baseline
// > database, table, and column names must be 255 characters or fewer
// > Athena accepts mixed case in DDL and DML queries, but lower cases the names when it executes the query
// > avoid using mixed case for table or column names
// > special characters other than underscore (_) are not supported
matchGlueName = regexp.MustCompile(`^[a-z0-9_]{1,255}$`).MatchString
// matchRolesAnywhereTrustAnchorName is a regex that matches against AWS IAM Roles Anywhere Trust Anchor Names.
// See https://docs.aws.amazon.com/rolesanywhere/latest/APIReference/API_CreateTrustAnchor.html#API_CreateTrustAnchor_RequestBody
matchRolesAnywhereTrustAnchorName = baseResourceNameMatcher
// matchRolesAnywhereProfileName is a regex that matches against AWS IAM Roles Anywhere Profile Names.
// See https://docs.aws.amazon.com/rolesanywhere/latest/APIReference/API_CreateProfile.html#API_CreateProfile_RequestBody
matchRolesAnywhereProfileName = baseResourceNameMatcher
baseResourceNameMatcher = regexp.MustCompile(`^[ a-zA-Z0-9-_]{1,255}$`).MatchString
// https://docs.aws.amazon.com/IAM/latest/UserGuide/reference-arns.html
validPartitions = []string{"aws", "aws-cn", "aws-us-gov"}
)
/*
Copyright 2022 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.
*/
package aws
// GetPartitionFromRegion get aws partition from region
// example, region "us-east-1" corresponds to partition "aws"
// region "cn-north-1" corresponds to partition "aws-cn"
func GetPartitionFromRegion(region string) string {
var partition string
switch {
case IsCNRegion(region):
partition = CNPartition
case IsUSGovRegion(region):
partition = USGovPartition
default:
partition = StandardPartition
}
return partition
}
const (
// StandardPartition is the partition ID of the AWS Standard partition.
StandardPartition = "aws"
// CNPartition is the partition ID of the AWS China partition.
CNPartition = "aws-cn"
// USGovPartition is the partition ID of the AWS GovCloud partition.
USGovPartition = "aws-us-gov"
)
/*
Copyright 2022 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.
*/
package aws
import (
"fmt"
"strconv"
"strings"
)
// IsCNRegion returns true if the region is an AWS China region.
func IsCNRegion(region string) bool {
return strings.HasPrefix(strings.ToLower(region), CNRegionPrefix)
}
// IsUSGovRegion returns true if the region is an AWS US GovCloud region.
func IsUSGovRegion(region string) bool {
return strings.HasPrefix(strings.ToLower(region), USGovRegionPrefix)
}
// ShortRegionToRegion converts short region codes to regular region names. For
// example, a short region "use1" maps to region "us-east-1".
//
// There is no official documentation on this mapping. Here is gist of others
// collecting these naming schemes:
// https://gist.github.com/colinvh/14e4b7fb6b66c29f79d3
//
// This function currently does not support regions in secert partitions.
func ShortRegionToRegion(shortRegion string) (string, bool) {
var prefix, direction string
// Determine region prefix.
remain := strings.ToLower(shortRegion)
switch {
case strings.HasPrefix(remain, "usg"):
prefix = USGovRegionPrefix
remain = remain[3:]
case strings.HasPrefix(remain, "cn"):
prefix = CNRegionPrefix
remain = remain[2:]
default:
// For regions in standard partition, the first two letters is the
// continent or country code (e.g. "eu" for Europe, "us" for US).
if len(remain) < 2 {
return "", false
}
prefix = remain[:2] + "-"
remain = remain[2:]
}
// Map direction codes.
switch {
case strings.HasPrefix(remain, "nw"):
direction = "northwest"
remain = remain[2:]
case strings.HasPrefix(remain, "ne"):
direction = "northeast"
remain = remain[2:]
case strings.HasPrefix(remain, "se"):
direction = "southeast"
remain = remain[2:]
case strings.HasPrefix(remain, "sw"):
direction = "southwest"
remain = remain[2:]
case strings.HasPrefix(remain, "n"):
direction = "north"
remain = remain[1:]
case strings.HasPrefix(remain, "e"):
direction = "east"
remain = remain[1:]
case strings.HasPrefix(remain, "w"):
direction = "west"
remain = remain[1:]
case strings.HasPrefix(remain, "s"):
direction = "south"
remain = remain[1:]
case strings.HasPrefix(remain, "c"):
direction = "central"
remain = remain[1:]
default:
return "", false
}
// Remain should be a number.
if _, err := strconv.Atoi(remain); err != nil {
return "", false
}
return fmt.Sprintf("%s%s-%s", prefix, direction, remain), true
}
const (
// CNRegionPrefix is the prefix for all AWS China regions.
CNRegionPrefix = "cn-"
// USGovRegionPrefix is the prefix for all AWS US GovCloud regions.
USGovRegionPrefix = "us-gov-"
)
/*
Copyright 2022 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.
*/
package azure
import (
"net"
"net/url"
"strings"
"github.com/gravitational/trace"
)
// IsAzureEndpoint returns true if the input URI is an Azure endpoint.
//
// The code implements approximate solution based on:
// - https://management.azure.com/metadata/endpoints?api-version=2019-05-01
// - https://github.com/Azure/azure-cli/blob/dev/src/azure-cli-core/azure/cli/core/cloud.py
func IsAzureEndpoint(hostname string) bool {
suffixes := []string{
"management.azure.com",
"graph.windows.net",
"batch.core.windows.net",
"rest.media.azure.net",
"datalake.azure.net",
"management.core.windows.net",
"gallery.azure.com",
"azuredatalakestore.net",
"azurecr.io",
"database.windows.net",
"azuredatalakeanalytics.net",
"vault.azure.net",
"core.windows.net",
"azurefd.net",
"login.microsoftonline.com", // required for "az logout"
"graph.microsoft.com", // Azure AD
}
for _, suffix := range suffixes {
// exact match
if hostname == suffix {
return true
}
// .suffix match
if strings.HasSuffix(hostname, "."+suffix) {
return true
}
}
return false
}
// IsDatabaseEndpoint returns true if provided endpoint is a valid database
// endpoint.
func IsDatabaseEndpoint(endpoint string) bool {
return strings.Contains(endpoint, DatabaseEndpointSuffix)
}
// IsCacheForRedisEndpoint returns true if provided endpoint is a valid Azure
// Cache for Redis endpoint.
func IsCacheForRedisEndpoint(endpoint string) bool {
return IsRedisEndpoint(endpoint) || IsRedisEnterpriseEndpoint(endpoint)
}
// IsRedisEndpoint returns true if provided endpoint is a valid Redis
// (non-Enterprise tier) endpoint.
func IsRedisEndpoint(endpoint string) bool {
return strings.Contains(endpoint, RedisEndpointSuffix)
}
// IsRedisEnterpriseEndpoint returns true if provided endpoint is a valid Redis
// Enterprise endpoint.
func IsRedisEnterpriseEndpoint(endpoint string) bool {
return strings.Contains(endpoint, RedisEnterpriseEndpointSuffix)
}
// IsMSSQLServerEndpoint returns true if provided endpoint is a valid SQL server
// database endpoint.
func IsMSSQLServerEndpoint(endpoint string) bool {
return strings.Contains(endpoint, MSSQLEndpointSuffix)
}
// ParseDatabaseEndpoint extracts database server name from Azure endpoint.
func ParseDatabaseEndpoint(endpoint string) (name string, err error) {
host, _, err := net.SplitHostPort(endpoint)
if err != nil {
return "", trace.Wrap(err)
}
// Azure endpoint looks like this:
// name.mysql.database.azure.com
parts := strings.Split(host, ".")
if !strings.HasSuffix(host, DatabaseEndpointSuffix) || len(parts) != 5 {
return "", trace.BadParameter("failed to parse %v as Azure endpoint", endpoint)
}
return parts[0], nil
}
// ParseCacheForRedisEndpoint extracts database server name from Azure Cache
// for Redis endpoint.
func ParseCacheForRedisEndpoint(endpoint string) (name string, err error) {
// Note that the Redis URI may contain schema and parameters.
host, err := GetHostFromRedisURI(endpoint)
if err != nil {
return "", trace.Wrap(err)
}
switch {
// Redis (non-Enterprise) endpoint looks like this:
// name.redis.cache.windows.net
case strings.HasSuffix(host, RedisEndpointSuffix):
return strings.TrimSuffix(host, RedisEndpointSuffix), nil
// Redis Enterprise endpoint looks like this:
// name.region.redisenterprise.cache.azure.net
case strings.HasSuffix(host, RedisEnterpriseEndpointSuffix):
name, _, ok := strings.Cut(strings.TrimSuffix(host, RedisEnterpriseEndpointSuffix), ".")
if !ok {
return "", trace.BadParameter("failed to parse %v as Azure Cache endpoint", endpoint)
}
return name, nil
default:
return "", trace.BadParameter("failed to parse %v as Azure Cache endpoint", endpoint)
}
}
// GetHostFromRedisURI extracts host name from a Redis URI. The URI may start
// with "redis://", "rediss://", or without. The URI may also have parameters
// like "?mode=cluster".
func GetHostFromRedisURI(uri string) (string, error) {
// Add a temporary schema to make a valid URL for url.Parse if schema is
// not found.
if !strings.Contains(uri, "://") {
uri = "schema://" + uri
}
parsed, err := url.Parse(uri)
if err != nil {
return "", trace.Wrap(err)
}
return parsed.Hostname(), nil
}
// ParseMSSQLEndpoint extracts database server name from Azure endpoint.
func ParseMSSQLEndpoint(endpoint string) (name string, err error) {
host, _, err := net.SplitHostPort(endpoint)
if err != nil {
return "", trace.Wrap(err)
}
// Azure endpoint looks like this:
// name.database.windows.net
parts := strings.Split(host, ".")
if !strings.HasSuffix(host, MSSQLEndpointSuffix) || len(parts) != 4 {
return "", trace.BadParameter("failed to parse %v as Azure MSSQL endpoint", endpoint)
}
if parts[0] == "" {
return "", trace.BadParameter("endpoint %v must contain database name", endpoint)
}
return parts[0], nil
}
const (
// DatabaseEndpointSuffix is the Azure database endpoint suffix. Used for
// MySQL, PostgreSQL, etc.
DatabaseEndpointSuffix = ".database.azure.com"
// RedisEndpointSuffix is the endpoint suffix for Redis.
RedisEndpointSuffix = ".redis.cache.windows.net"
// RedisEnterpriseEndpointSuffix is the endpoint suffix for Redis Enterprise.
RedisEnterpriseEndpointSuffix = ".redisenterprise.cache.azure.net"
// MSSQLEndpointSuffix is the Azure SQL Server endpoint suffix.
MSSQLEndpointSuffix = ".database.windows.net"
)
/*
Copyright 2022 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.
*/
package azure
import "strings"
// NormalizeLocation converts a Azure location in various formats to the same
// simple format.
//
// This function assumes the input location is in one of the following formats:
// - Name (the "simple" format): "northcentralusstage"
// - Display name: "North Central US (Stage)"
// - Regional display name: "(US) North Central US (Stage)"
//
// Note that the location list can be generated from `az account list-locations
// -o table`. However, this CLI command only lists the locations for the
// current active subscription so it may not show locations in other
// parititions like Government or China.
func NormalizeLocation(input string) string {
if input == "" {
return input
}
// Check if the input is a recognized simple name.
if _, found := locationsToDisplayNames[input]; found {
return input
}
// If input starts with '(', it should be the Regional display name. Then
// removes the first bracket and its content. The leftover will be a
// display name.
if input[0] == '(' {
if index := strings.IndexRune(input, ')'); index >= 0 {
input = input[index:]
}
input = strings.TrimSpace(input)
}
// Check if the input is a recognized display name. If so, return the
// simple name from the mapping.
if location, found := displayNamesToLocations[input]; found {
return location
}
// Try our best to convert an unregconized input:
// - Remove brackets and spaces.
// - To lower case.
replacer := strings.NewReplacer("(", "", ")", "", " ", "")
return strings.ToLower(replacer.Replace(input))
}
var (
// displayNamesToLocations maps a location's "Display Name" to its simple
// "Name".
displayNamesToLocations = map[string]string{
// Azure locations.
"East US": "eastus",
"East US 2": "eastus2",
"South Central US": "southcentralus",
"West US 2": "westus2",
"West US 3": "westus3",
"Australia East": "australiaeast",
"Southeast Asia": "southeastasia",
"North Europe": "northeurope",
"Sweden Central": "swedencentral",
"UK South": "uksouth",
"West Europe": "westeurope",
"Central US": "centralus",
"South Africa North": "southafricanorth",
"Central India": "centralindia",
"East Asia": "eastasia",
"Japan East": "japaneast",
"Korea Central": "koreacentral",
"Canada Central": "canadacentral",
"France Central": "francecentral",
"Germany West Central": "germanywestcentral",
"Norway East": "norwayeast",
"Switzerland North": "switzerlandnorth",
"UAE North": "uaenorth",
"Brazil South": "brazilsouth",
"East US 2 EUAP": "eastus2euap",
"Qatar Central": "qatarcentral",
"Central US (Stage)": "centralusstage",
"East US (Stage)": "eastusstage",
"East US 2 (Stage)": "eastus2stage",
"North Central US (Stage)": "northcentralusstage",
"South Central US (Stage)": "southcentralusstage",
"West US (Stage)": "westusstage",
"West US 2 (Stage)": "westus2stage",
"Asia": "asia",
"Asia Pacific": "asiapacific",
"Australia": "australia",
"Brazil": "brazil",
"Canada": "canada",
"Europe": "europe",
"France": "france",
"Germany": "germany",
"Global": "global",
"India": "india",
"Japan": "japan",
"Korea": "korea",
"Norway": "norway",
"Singapore": "singapore",
"South Africa": "southafrica",
"Switzerland": "switzerland",
"United Arab Emirates": "uae",
"United Kingdom": "uk",
"United States": "unitedstates",
"United States EUAP": "unitedstateseuap",
"East Asia (Stage)": "eastasiastage",
"Southeast Asia (Stage)": "southeastasiastage",
"East US STG": "eastusstg",
"South Central US STG": "southcentralusstg",
"North Central US": "northcentralus",
"West US": "westus",
"Jio India West": "jioindiawest",
"Central US EUAP": "centraluseuap",
"West Central US": "westcentralus",
"South Africa West": "southafricawest",
"Australia Central": "australiacentral",
"Australia Central 2": "australiacentral2",
"Australia Southeast": "australiasoutheast",
"Japan West": "japanwest",
"Jio India Central": "jioindiacentral",
"Korea South": "koreasouth",
"South India": "southindia",
"West India": "westindia",
"Canada East": "canadaeast",
"France South": "francesouth",
"Germany North": "germanynorth",
"Norway West": "norwaywest",
"Switzerland West": "switzerlandwest",
"UK West": "ukwest",
"UAE Central": "uaecentral",
"Brazil Southeast": "brazilsoutheast",
// Azure Government locations.
//
// https://learn.microsoft.com/en-us/azure/azure-government/documentation-government-get-started-connect-with-ps
"USDoD Central": "usdodcentral",
"USDoD East": "usdodeast",
"USGov Arizona": "usgovarizona",
"USGov Iowa": "usgoviowa",
"USGov Texas": "usgovtexas",
"USGov Virginia": "usgovvirginia",
"USSec East": "usseceast",
"USSec West": "ussecwest",
"USSec West Central": "ussecwestcentral",
// Azure China locations.
"China East": "chinaeast",
"China East 2": "chinaeast2",
"China North": "chinanorth",
"China North 2": "chinanorth2",
"China North 3": "chinanorth3",
}
// locationsToDisplayNames maps Azure location names to their display
// names. This is the reverse lookup map of displayNamesToLocations.
locationsToDisplayNames = map[string]string{}
)
func initLocations() {
for displayName, location := range displayNamesToLocations {
locationsToDisplayNames[location] = displayName
}
}
func init() {
initLocations()
}
/*
* Teleport
* Copyright (C) 2026 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 appresource
import "encoding/json"
// Hint explains a near-miss on one [Rule], where its path and method matched but its
// Where did not. It contains the rule's DenyCodeHint and DenyReasonHint.
type Hint struct {
Code string `json:"code"`
Reason string `json:"reason,omitempty"`
}
// DenyKind is the category of a denial, emitted as deny_kind on the
// app.session.request.denied audit event.
type DenyKind string
const (
// DenyNotAllowed is the kind for a well-formed request that no allow
// rule matched.
DenyNotAllowed DenyKind = "teleport_request_not_allowed"
// DenyRoleVersionUnsupported is the kind for a request denied because a
// role or a rule is a version this agent cannot evaluate, as in a
// mixed-version cluster.
DenyRoleVersionUnsupported DenyKind = "teleport_role_version_unsupported"
// DenyInvalidRequest is the denial category for a malformed request
// path, e.g. containing a ".." segment. No rule is evaluated.
DenyInvalidRequest DenyKind = "teleport_invalid_request"
)
// Decision is the aggregated result of evaluating one request against the
// app_resources rules on the caller's roles.
type Decision struct {
// Allowed is true if any rule matched.
Allowed bool
// Allow contains details iff Allowed is true.
Allow *AllowDetails
// Deny contains details iff Allowed is false.
Deny *DenyDetails
// EvaluatedRoles lists the roles evaluated, in evaluation order.
EvaluatedRoles []string
}
// AllowDetails is an allow decision record derived from the matching rule.
type AllowDetails struct {
// Vars contains the path segments the matching rule captured.
Vars map[string]string
// Code is the matching rule's allow_code.
Code string
// Reason is the matching rule's allow_reason.
Reason string
}
// DenyDetails is a deny decision record.
type DenyDetails struct {
// Kind is the structured reason for the deny.
Kind DenyKind
// Hints lists every hint that fired, in rule order.
Hints []Hint
}
// decisionJSON is the flat wire form of a Decision.
type decisionJSON struct {
Allowed bool `json:"allowed"`
EvaluatedRoles []string `json:"evaluated_roles,omitempty"`
Vars map[string]string `json:"vars,omitempty"`
AllowCode string `json:"allow_code,omitempty"`
AllowReason string `json:"allow_reason,omitempty"`
DenyKind DenyKind `json:"deny_kind,omitempty"`
Hints []Hint `json:"hints,omitempty"`
}
// MarshalJSON encodes the decision in its flat wire form, with unset fields
// omitted. Detail on the side that does not match Allowed is dropped.
func (d Decision) MarshalJSON() ([]byte, error) {
out := decisionJSON{
Allowed: d.Allowed,
EvaluatedRoles: d.EvaluatedRoles,
}
if d.Allowed && d.Allow != nil {
out.Vars = d.Allow.Vars
out.AllowCode = d.Allow.Code
out.AllowReason = d.Allow.Reason
}
if !d.Allowed && d.Deny != nil {
out.DenyKind = d.Deny.Kind
out.Hints = d.Deny.Hints
}
return json.Marshal(out)
}
/*
* Teleport
* Copyright (C) 2026 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 appresource
import (
"slices"
"strings"
"github.com/gravitational/trace"
)
// maxWhereBytes is the maximum length in bytes of one where clause, the sugared
// form.
const maxWhereBytes = 1 << 10 // 1 KiB
// maxReasonBytes is the maximum length in bytes of an allow_reason or
// deny_reason_hint.
const maxReasonBytes = 1 << 10 // 1 KiB
// maxAuditCodeBytes is the maximum length in bytes of an allow_code or
// deny_code_hint.
const maxAuditCodeBytes = 256
// maxPathBytes is the maximum length in bytes of one path pattern.
const maxPathBytes = 1 << 10 // 1 KiB
// maxPaths is the maximum number of path patterns in one rule.
const maxPaths = 64
// Rule is one app_resources entry, the sugared form. A request matches when its
// path matches Paths, its method matches Methods, and its Where clause evaluates to
// true.
type Rule struct {
// Paths are the path patterns the rule matches. The {project} segment in
// "/api/projects/{project}/**" is captured, and Where reads it as
// vars.project. A rule sets either Paths or AllowAll.
Paths []string `yaml:"paths,omitempty"`
// Methods is a list of GET, HEAD, POST, PUT, PATCH, DELETE, OPTIONS, or
// TRACE, matched case-insensitively. A request method is not folded, so
// it must be upper case. Unset, Methods allows all eight.
Methods []string `yaml:"methods,omitempty"`
// Where is a predicate over the caller identity and the rule's path
// captures, such as contains(user.traits["projects"], vars.project). If
// set, it must evaluate to true for the rule to match.
Where string `yaml:"where,omitempty"`
// AllowEncoded lists the characters a request path may carry in
// percent-encoded form for the rule to match. The only supported value is
// "/", which allows the encoded slash, %2F or %2f.
AllowEncoded []string `yaml:"allow_encoded,omitempty"`
// AllowCode is the code recorded on the allow audit event when the rule
// matches. If it is not set, no allow audit event is recorded. A code may
// not start with the reserved "teleport_" prefix.
AllowCode string `yaml:"allow_code,omitempty"`
// AllowReason is the explanation recorded alongside AllowCode. A rule sets
// it only together with AllowCode.
AllowReason string `yaml:"allow_reason,omitempty"`
// DenyCodeHint is the code added to the deny decision when the rule's path
// and method match but the Where predicate does not. A denied request
// collects a code from every such rule, so one decision can record several
// codes. A code may not start with the reserved "teleport_" prefix.
DenyCodeHint string `yaml:"deny_code_hint,omitempty"`
// DenyReasonHint is the explanation recorded alongside DenyCodeHint. A rule
// sets it only together with DenyCodeHint.
DenyReasonHint string `yaml:"deny_reason_hint,omitempty"`
// AllowAll grants unrestricted access to every path and method. It cannot
// be combined with any other field.
AllowAll bool `yaml:"allow_all,omitempty"`
}
// validateAuditCode checks an allow or deny code. A valid code is 1 to 256 bytes of
// [a-z0-9_] and does not start with the reserved teleport_ prefix.
func validateAuditCode(code string) error {
if len(code) < 1 || len(code) > maxAuditCodeBytes {
return trace.BadParameter("code %q must be 1 to %d bytes", code, maxAuditCodeBytes)
}
for _, r := range code {
legal := (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9') || r == '_'
if !legal {
return trace.BadParameter("code %q must contain only [a-z0-9_]", code)
}
}
if strings.HasPrefix(code, "teleport_") {
return trace.BadParameter("code %q must not start with the reserved teleport_ prefix", code)
}
return nil
}
// validate checks a rule's structural constraints, e.g. that AllowAll cannot be
// combined with another field. Path pattern checks are left to compile time.
func (r Rule) validate() error {
if r.AllowAll {
return r.validateAllowAllStandsAlone()
}
if len(r.Paths) == 0 {
return trace.BadParameter("a rule must set paths or allow_all")
}
if err := validatePaths(r.Paths); err != nil {
return trace.Wrap(err)
}
if err := validateMethods(r.Methods); err != nil {
return trace.Wrap(err)
}
if err := validateWhere(r.Where); err != nil {
return trace.Wrap(err)
}
for _, e := range r.AllowEncoded {
if e != "/" {
return trace.BadParameter("allow_encoded allows only the separator %q, got %q", "/", e)
}
}
if r.AllowReason != "" && r.AllowCode == "" {
return trace.BadParameter("allow_reason set without allow_code")
}
if r.AllowCode != "" {
if err := validateAuditCode(r.AllowCode); err != nil {
return trace.Wrap(err, "invalid allow_code")
}
}
if len(r.AllowReason) > maxReasonBytes {
return trace.BadParameter("allow_reason is %d bytes, over the %d byte cap", len(r.AllowReason), maxReasonBytes)
}
if r.DenyReasonHint != "" && r.DenyCodeHint == "" {
return trace.BadParameter("deny_reason_hint set without deny_code_hint")
}
if r.DenyCodeHint != "" {
if err := validateAuditCode(r.DenyCodeHint); err != nil {
return trace.Wrap(err, "invalid deny_code_hint")
}
if r.Where == "" {
return trace.BadParameter("deny_code_hint set without a where clause")
}
}
if len(r.DenyReasonHint) > maxReasonBytes {
return trace.BadParameter("deny_reason_hint is %d bytes, over the %d byte cap", len(r.DenyReasonHint), maxReasonBytes)
}
return nil
}
// validateAllowAllStandsAlone rejects an allow_all rule that also sets another
// field.
func (r Rule) validateAllowAllStandsAlone() error {
if len(r.Paths) > 0 || len(r.Methods) > 0 || r.Where != "" ||
len(r.AllowEncoded) > 0 || r.AllowCode != "" || r.AllowReason != "" ||
r.DenyCodeHint != "" || r.DenyReasonHint != "" {
return trace.BadParameter("allow_all cannot be combined with any other field")
}
return nil
}
// validatePaths checks the count and byte caps on a rule's path patterns. The
// pattern syntax is checked when the rule compiles.
func validatePaths(paths []string) error {
if len(paths) > maxPaths {
return trace.BadParameter("a rule holds %d paths, over the cap of %d", len(paths), maxPaths)
}
for _, p := range paths {
if len(p) > maxPathBytes {
return trace.BadParameter("path is %d bytes, over the %d byte cap", len(p), maxPathBytes)
}
}
return nil
}
// validateMethods rejects a name outside validMethods, folded to upper case, so
// "get" passes and "GTE" fails.
func validateMethods(methods []string) error {
for _, m := range methods {
if !slices.Contains(validMethods, strings.ToUpper(m)) {
return trace.BadParameter("method %q is not one of %s", m, strings.Join(validMethods, ", "))
}
}
return nil
}
// validateWhere checks the byte cap on a sugared rule's where clause. The where
// language itself is checked when the clause compiles.
func validateWhere(where string) error {
if where == "" {
return nil
}
if len(where) > maxWhereBytes {
return trace.BadParameter("where clause is %d bytes, over the %d byte cap", len(where), maxWhereBytes)
}
return nil
}
/*
* Teleport
* Copyright (C) 2026 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 appresource checks whether an HTTP app request is allowed
// by a role. Roles carry allow-only rules. A rule can match on
// request path, HTTP method, and a where predicate over the user
// identity. A rule sets either paths, with the other fields
// optional, or allow_all, which stands alone.
//
// Example role fragment:
//
// allow:
// app_resources:
// - paths:
// - /api/v4/user/{username}
// where: user.name == vars.username
package appresource
import (
"strconv"
"strings"
"unicode"
"unicode/utf8"
"github.com/gravitational/trace"
"golang.org/x/text/unicode/norm"
)
// maxPathLength bounds the length of a path Tokenize accepts.
const maxPathLength = 8 << 10 // 8 KiB
// legalPathPunct is the non-alphanumeric bytes allowed in a raw
// URL path. It is RFC 3986 pchar except for ";", plus "/" and "%".
//
// ";" is dropped because matrix parameters and ";jsessionid" may
// cause the matcher and the upstream app to disagree on where the
// path ends.
const legalPathPunct = "-._~!$&'()*+,=:@/%"
// Tokenize validates an HTTP request path and splits it on a real
// "/" into the encoded segments a role's path rules match against, so
// an encoded slash stays inside one segment. Pass
// [net/url.URL.EscapedPath], the encoded path sent to the upstream
// app, not the already-decoded [net/url.URL.Path].
//
// Tokenize accepts a path that starts with "/", stays under 8 KiB,
// and holds only the path characters RFC 3986 allows, except for ";".
// Anything else has to be sent percent-encoded, and the only escapes
// allowed are the separator %2F ("/"), the space %20 (" "), and the
// UTF-8 bytes of non-ASCII text. Tokenize also rejects a path an
// upstream app could read as a different path than the one a role
// matched, such as "/a/../b" or "/files/secret." on a server that
// trims a trailing dot.
func Tokenize(path string) ([]string, error) {
if len(path) > maxPathLength {
return nil, trace.BadParameter("path length %d exceeds the %d byte limit", len(path), maxPathLength)
}
if !strings.HasPrefix(path, "/") {
return nil, trace.BadParameter("path %q must start with /", clip(path))
}
if err := validateRawBytes(path); err != nil {
return nil, trace.Wrap(err)
}
if err := validateDecoded(path); err != nil {
return nil, trace.Wrap(err)
}
return strings.Split(path[1:], "/"), nil
}
// validateRawBytes rejects any byte that cannot appear in a URL path
// under RFC 3986, any invalid percent-escape, and every escape except
// %2F ("/"), %20 (" "), or one that decodes to a non-ASCII byte.
func validateRawBytes(path string) error {
for i := 0; i < len(path); i++ {
if !isLegalPathByte(path[i]) {
return trace.BadParameter("path %q contains an illegal URL byte %q", clip(path), path[i:i+1])
}
if path[i] != '%' {
continue
}
if i+2 >= len(path) {
return trace.BadParameter("path %q has a truncated percent-escape", clip(path))
}
v, err := strconv.ParseUint(path[i+1:i+3], 16, 8)
if err != nil {
return trace.BadParameter("path %q has a malformed percent-escape %q", clip(path), path[i:i+3])
}
if !isAllowedEscape(byte(v)) {
const msg = "path %q contains the percent-escape %q; only the encoded separator %%2F, the encoded space %%20, and non-ASCII content escapes are allowed"
return trace.BadParameter(msg, clip(path), path[i:i+3])
}
i += 2
}
return nil
}
// isLegalPathByte reports whether a given byte may appear in a raw
// URL path.
func isLegalPathByte(b byte) bool {
switch {
case b >= 'A' && b <= 'Z', b >= 'a' && b <= 'z', b >= '0' && b <= '9':
return true
}
return strings.IndexByte(legalPathPunct, b) >= 0
}
// isAllowedEscape reports whether a percent-escape decoding to b may
// appear in a path. [validateRawBytes] accepts exactly the escapes
// [decode] resolves, which are the separator, the space, and
// non-ASCII content.
func isAllowedEscape(b byte) bool {
return b == '/' || b == ' ' || b >= 0x80
}
// validateDecoded checks the decoded validation view of the path. It
// rejects consecutive slashes, "." and ".." segments, and any content
// that is not NFKC-stable graphic UTF-8.
func validateDecoded(path string) error {
decoded := decode(path)
if strings.Contains(decoded, "//") {
const msg = "path %q has consecutive slashes once the encoded separator %%2F is decoded"
return trace.BadParameter(msg, clip(path))
}
if !utf8.ValidString(decoded) {
return trace.BadParameter("path %q is not valid UTF-8 once decoded", clip(path))
}
for seg := range strings.SplitSeq(decoded[1:], "/") {
if err := rejectDotSegment(seg); err != nil {
return trace.Wrap(err)
}
if err := rejectLeadingMark(seg); err != nil {
return trace.Wrap(err)
}
if err := rejectEdgeSpace(seg); err != nil {
return trace.Wrap(err)
}
if err := rejectTrailingDot(seg); err != nil {
return trace.Wrap(err)
}
}
if !norm.NFKC.IsNormalString(decoded) {
return trace.BadParameter("path %q is not NFKC-normalized", clip(path))
}
for _, r := range decoded {
if !isGraphicRune(r) {
const msg = "path %q contains the disallowed character %q; only letters, marks, numbers, punctuation, symbols, and the encoded space %%20 are allowed"
return trace.BadParameter(msg, clip(path), string(r))
}
}
return nil
}
// decode returns the decoded validation view of path s, resolving
// only valid escapes. %2F and %2f become "/", %20 a space, and a
// non-ASCII escape its byte. The view exposes a structural byte
// written as an escape, so "/x%2F..%2Fadmin" is rejected for the same
// reason ".." is rejected in "/x/../admin".
func decode(s string) string {
if !strings.ContainsRune(s, '%') {
return s
}
var b strings.Builder
b.Grow(len(s))
for i := 0; i < len(s); i++ {
if s[i] == '%' && i+2 < len(s) {
if v, err := strconv.ParseUint(s[i+1:i+3], 16, 8); err == nil && isAllowedEscape(byte(v)) {
b.WriteByte(byte(v))
i += 2
continue
}
}
b.WriteByte(s[i])
}
return b.String()
}
// rejectDotSegment rejects a segment made of only dots and spaces
// that has at least one dot. "." and ".." are traversal segments,
// and an upstream that strips spaces or extra dots would resolve
// forms like ". ." or "..." the same way.
func rejectDotSegment(seg string) error {
if strings.Trim(seg, ". ") == "" && strings.Contains(seg, ".") {
const msg = `segment %q is only dots and spaces; an upstream could resolve it as "." or ".."`
return trace.BadParameter(msg, clip(seg))
}
return nil
}
// rejectLeadingMark rejects a segment whose first character composes
// onto the character before it, so "/a/%CC%87b", where %CC%87 is
// U+0307 combining dot above, looks like "/a/b". RFC 5891 bans the
// same form at the start of a domain label.
func rejectLeadingMark(seg string) error {
if seg == "" || seg[0] < utf8.RuneSelf {
return nil
}
r, _ := utf8.DecodeRuneInString(seg)
// Some composing characters are not marks, so both checks are needed.
if unicode.IsMark(r) || !norm.NFKC.PropertiesString(seg).BoundaryBefore() {
const msg = "segment %q starts with %q, which composes onto the character before it; it must follow a base character"
return trace.BadParameter(msg, clip(seg), string(r))
}
return nil
}
// rejectEdgeSpace rejects a segment whose first or last rune is a
// space. An upstream that trims the segment would see a different
// one, so "..%20" would trim to ".." and "secret%20" to "secret".
func rejectEdgeSpace(seg string) error {
if strings.HasPrefix(seg, " ") || strings.HasSuffix(seg, " ") {
const msg = "segment %q starts or ends with a space; a space must be between other characters"
return trace.BadParameter(msg, clip(seg))
}
return nil
}
// rejectTrailingDot rejects a segment whose trailing run of dots and
// spaces contains a dot. IIS and Windows trim trailing dots and
// spaces, so an upstream would resolve "secret." and "secret%20." as
// "secret", a segment the matcher never saw. A leading dot stays
// allowed because paths such as "/.well-known" depend on it.
func rejectTrailingDot(seg string) error {
trimmed := strings.TrimRight(seg, ". ")
if strings.Contains(seg[len(trimmed):], ".") {
const msg = "segment %q ends with dots and spaces; an upstream could trim it to %q"
return trace.BadParameter(msg, clip(seg), clip(trimmed))
}
return nil
}
// isGraphicRune reports whether r is U+0020, a letter, mark, number,
// punctuation, or symbol. U+0020 is allowed because it only reaches
// the decoded view as %20, and [rejectEdgeSpace] keeps it off the
// segment edges. Every other space and separator rune that
// unicode.IsGraphic allows is excluded, along with control, format,
// surrogate, private-use, and unassigned runes. This is the only
// check that rejects the NFKC-stable ones among them, such as U+1680
// ogham space mark.
func isGraphicRune(r rune) bool {
return r == ' ' || unicode.IsLetter(r) || unicode.IsMark(r) ||
unicode.IsNumber(r) || unicode.IsPunct(r) || unicode.IsSymbol(r)
}
// clip shortens s for use in an error message. A rejected path can be
// kilobytes long, and the message survives into logs and audit events.
func clip(s string) string {
const limit = 256
if len(s) <= limit {
return s
}
return s[:limit] + "..."
}
/*
* Teleport
* Copyright (C) 2026 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 appresource
import (
"net/http"
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/utils/typical"
)
// validMethods are the HTTP methods a where clause evaluation accepts.
var validMethods = []string{
http.MethodGet,
http.MethodHead,
http.MethodPost,
http.MethodPut,
http.MethodPatch,
http.MethodDelete,
http.MethodOptions,
http.MethodTrace,
}
// whereParser is the shared cached parser for where clauses.
var whereParser = mustNewWhereParser()
// Request encodes the elements of the HTTP request a where clause is
// evaluated against.
type Request struct {
Method string
}
// Identity is the caller a where clause is evaluated against.
type Identity struct {
Name string
Roles []string
Traits map[string][]string
}
// Env holds the values one where clause evaluation reads. Request and
// Identity are deliberately not [http.Request] and tlsca.Identity, to make
// clear which fields matter for evaluation.
type Env struct {
Request Request
Identity Identity
}
// Where is a compiled where clause. Only CompileWhere returns a usable
// value. Evaluation writes nothing back to a Where, so a single Where can
// serve concurrent requests. The caller must not mutate the slices or
// map in Env during an evaluation.
type Where struct {
expression typical.Expression[Env, bool]
}
// NewEnv builds the environment a where clause is evaluated against and
// rejects a request no rule may authorize. Callers build it once per
// request, before matching any rule, because a rule without a where
// clause never reaches Evaluate.
func NewEnv(request Request, identity Identity) (Env, error) {
if err := validateMethod(request.Method); err != nil {
return Env{}, trace.Wrap(err)
}
return Env{Request: request, Identity: identity}, nil
}
// CompileWhere parses and type-checks a where clause.
func CompileWhere(expr string) (*Where, error) {
expression, err := whereParser.Parse(expr)
if err != nil {
// The aggregate classifies the result as BadParameter for the
// caller and keeps typical's typed error for errors.As. Parse does
// not return a BadParameter for every failure.
return nil, trace.NewAggregate(trace.BadParameter("compiling where clause %q", expr), err)
}
return &Where{expression: expression}, nil
}
// Evaluate reports whether the where clause matches the environment. The
// result is only meaningful when the error is nil.
func (w *Where) Evaluate(env Env) (bool, error) {
if err := validateMethod(env.Request.Method); err != nil {
return false, trace.Wrap(err)
}
match, err := w.expression.Evaluate(env)
if err != nil {
return false, trace.Wrap(err)
}
return match, nil
}
// validateMethod rejects a request method outside the canonical HTTP
// method list. NewEnv runs it at the request boundary and Evaluate
// repeats it, so an environment a caller assembled itself still cannot
// authorize such a request.
func validateMethod(method string) error {
if !slices.Contains(validMethods, method) {
return trace.BadParameter("unsupported HTTP method %q", method)
}
return nil
}
// mustNewWhereParser builds the where clause parser and panics if the
// parser spec is invalid or the expression cache cannot be built.
func mustNewWhereParser() *typical.CachedParser[Env, bool] {
p, err := typical.NewCachedParser[Env, bool](typical.ParserSpec[Env]{
Variables: map[string]typical.Variable{
// true and false are bound because typical has no bool literal.
"true": true,
"false": false,
"user.name": typical.DynamicVariable(func(e Env) (string, error) {
return e.Identity.Name, nil
}),
"user.roles": typical.DynamicVariable(func(e Env) ([]string, error) {
return e.Identity.Roles, nil
}),
// A key the identity does not have reads as an empty list, as in
// the role where-clause language, so a mistyped key under a
// negation matches every caller.
"user.traits": typical.DynamicMapFunction(func(e Env, key string) ([]string, error) {
return e.Identity.Traits[key], nil
}),
"request.method": typical.DynamicVariable(func(e Env) (string, error) {
return e.Request.Method, nil
}),
},
Functions: map[string]typical.Function{
// set and contains are named after the functions in the role
// where-clause language, services.NewWhereParser.
"set": typical.UnaryVariadicFunction[Env](func(args ...string) ([]string, error) {
return args, nil
}),
"contains": typical.BinaryFunction[Env](func(list []string, item string) (bool, error) {
return slices.Contains(list, item), nil
}),
"lower": typical.UnaryFunction[Env](func(s string) (string, error) {
return strings.ToLower(s), nil
}),
"upper": typical.UnaryFunction[Env](func(s string) (string, error) {
return strings.ToUpper(s), nil
}),
"has_prefix": typical.BinaryFunction[Env](func(s, prefix string) (bool, error) {
return strings.HasPrefix(s, prefix), nil
}),
"has_suffix": typical.BinaryFunction[Env](func(s, suffix string) (bool, error) {
return strings.HasSuffix(s, suffix), nil
}),
"has_substring": typical.BinaryFunction[Env](func(s, substr string) (bool, error) {
return strings.Contains(s, substr), nil
}),
},
})
if err != nil {
panic(trace.Wrap(err, "building the where clause parser (this is a bug)"))
}
return p
}
/*
* 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 webauthn
import (
"context"
"crypto/x509"
"encoding/pem"
"errors"
"log/slog"
"slices"
"github.com/go-webauthn/webauthn/protocol"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
var log = logutils.NewPackageLogger(teleport.ComponentKey, "WebAuthn")
// x5cFormats enumerates all attestation formats that supply an attestation
// chain through the "x5c" field.
// See https://www.w3.org/TR/webauthn/#sctn-defined-attestation-formats.
var x5cFormats = []string{
"packed",
"tpm",
"android-key",
"fido-u2f",
"apple",
}
func verifyAttestation(cfg *types.Webauthn, obj protocol.AttestationObject) error {
if len(cfg.AttestationAllowedCAs) == 0 && len(cfg.AttestationDeniedCAs) == 0 {
return nil // Attestation disabled.
}
attestationChain, err := getChainFromObj(obj)
if err != nil {
return trace.Wrap(
err, "failed to read attestation certificate; make sure you are using a device from a trusted manufacturer")
}
// We don't really expect errors at this stage, by the time the configuration
// gets here it was already validated by Teleport.
allowedPool, err := x509PEMsToCertPool(cfg.AttestationAllowedCAs)
if err != nil {
return trace.Wrap(err, "invalid webauthn attestation_allowed_ca")
}
deniedPool, err := x509PEMsToCertPool(cfg.AttestationDeniedCAs)
if err != nil {
return trace.Wrap(err, "invalid webauthn attestation_denied_ca")
}
verifyOptsBase := x509.VerifyOptions{
// TPM-bound certificates, like those issued for Windows Hello, set
// ExtKeyUsage OID 2.23.133.8.3, aka "AIK (Attestation Identity Key)
// certificate".
//
// There isn't an ExtKeyUsage constant for that, so we allow any.
//
// - https://learn.microsoft.com/en-us/windows/apps/develop/security/windows-hello#attestation
// - https://oid-base.com/get/2.23.133.8.3
KeyUsages: []x509.ExtKeyUsage{x509.ExtKeyUsageAny},
}
// Attestation check works as follows:
// 1. At least one certificate must belong to the allowed pool.
// 2. No certificates may belong to the denied pool.
//
// It is possible for both allowed and denied CAs to be present. It's also
// possible for configurations to allow a broad range of options (eg, all
// YubiKey devices) while denying a smaller subset (a certain model or lot),
// so both checks (allowed and denied) may be true for the same cert.
allowed := len(cfg.AttestationAllowedCAs) == 0
for _, cert := range attestationChain {
opts := verifyOptsBase // take copy
opts.Roots = allowedPool
if _, err := cert.Verify(opts); err == nil {
allowed = true // OK, but keep checking
} else {
log.DebugContext(context.Background(),
"Attestation check for allowed CAs failed",
"subject", cert.Subject,
"error", err,
)
}
opts = verifyOptsBase // take copy
opts.Roots = deniedPool
if _, err := cert.Verify(opts); err == nil {
return trace.BadParameter("attestation certificate %q from issuer %q not allowed", cert.Subject, cert.Issuer)
} else if !errors.As(err, new(x509.UnknownAuthorityError)) {
log.DebugContext(context.Background(),
"Attestation check for denied CAs failed",
"subject", cert.Subject,
"error", err,
)
}
}
if !allowed {
return trace.BadParameter(
"failed to verify device attestation certificate; make sure you are using a device from a trusted manufacturer")
}
return nil
}
func x509PEMsToCertPool(certPEMs []string) (*x509.CertPool, error) {
pool := x509.NewCertPool()
for _, cert := range certPEMs {
if !pool.AppendCertsFromPEM([]byte(cert)) {
return nil, trace.BadParameter("failed to parse certificate PEM")
}
}
return pool, nil
}
func getChainFromObj(obj protocol.AttestationObject) ([]*x509.Certificate, error) {
if slices.Contains(x5cFormats, obj.Format) {
return getChainFromX5C(obj)
}
if obj.Format == "none" {
// Return a nicer error for "none", since we do allow it in non-attestation
// scenarios.
return nil, trace.BadParameter("attestation format %q not allowed for direct attestation", obj.Format)
}
return nil, trace.BadParameter("attestation format %q not supported", obj.Format)
}
func getChainFromX5C(obj protocol.AttestationObject) ([]*x509.Certificate, error) {
x5c, ok := obj.AttStatement["x5c"]
if !ok {
// Warn about self-attestation and Touch ID, it may save someone some grief.
return nil, trace.BadParameter(
"%q attestation: self attestation not allowed; includes Touch ID in non-Apple browsers", obj.Format)
}
x5cArray, ok := x5c.([]any)
if !ok {
return nil, trace.BadParameter("%q attestation: unexpected x5c type: %T", obj.Format, x5c)
}
if len(x5cArray) == 0 {
return nil, trace.BadParameter("%q attestation: empty certificate chain", obj.Format)
}
chain := make([]*x509.Certificate, len(x5cArray))
for i, val := range x5cArray {
cert, ok := val.([]byte)
if !ok {
return nil, trace.BadParameter("%q attestation: unexpected x5c element type at index %v: %T", obj.Format, i, val)
}
var err error
chain[i], err = x509.ParseCertificate(cert)
if err != nil {
return nil, trace.Wrap(err, "%q attestation: failed to parse certificate at index %v", obj.Format, i)
}
}
// Print out attestation certs if debug is enabled.
// This may come in handy for people having trouble with their setups.
ctx := context.Background()
if log.Handler().Enabled(ctx, slog.LevelDebug) {
for _, cert := range chain {
certPEM := pem.EncodeToMemory(&pem.Block{
Type: "CERTIFICATE",
Bytes: cert.Raw,
})
log.DebugContext(context.Background(), "got attestation certificate",
"format", obj.Format,
"certificate", string(certPEM),
)
}
}
return chain, nil
}
/*
* 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 webauthn
import (
"github.com/go-webauthn/webauthn/protocol"
wan "github.com/go-webauthn/webauthn/webauthn"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/defaults"
)
const (
defaultDisplayName = "Teleport"
)
// webAuthnParams groups the parameters necessary for the creation of
// wan.WebAuthn instances.
type webAuthnParams struct {
cfg *types.Webauthn
rpID string
origin string
requireResidentKey bool
requireUserVerification bool
}
func newWebAuthn(p webAuthnParams) (*wan.WebAuthn, error) {
attestation := protocol.PreferNoAttestation
if len(p.cfg.AttestationAllowedCAs) > 0 || len(p.cfg.AttestationDeniedCAs) > 0 {
attestation = protocol.PreferDirectAttestation
}
residentKeyRequirement := protocol.ResidentKeyRequirementDiscouraged
if p.requireResidentKey {
residentKeyRequirement = protocol.ResidentKeyRequirementRequired
}
// Default to "discouraged", otherwise some browsers may do needless PIN
// prompts.
userVerification := protocol.VerificationDiscouraged
if p.requireUserVerification {
userVerification = protocol.VerificationRequired
}
timeoutConfig := wan.TimeoutConfig{
Enforce: true,
Timeout: defaults.WebauthnChallengeTimeout,
TimeoutUVD: defaults.WebauthnChallengeTimeout,
}
return wan.New(&wan.Config{
RPID: p.rpID,
RPOrigins: []string{p.origin},
RPDisplayName: defaultDisplayName,
AttestationPreference: attestation,
AuthenticatorSelection: protocol.AuthenticatorSelection{
RequireResidentKey: &p.requireResidentKey,
ResidentKey: residentKeyRequirement,
UserVerification: userVerification,
},
Timeouts: wan.TimeoutsConfig{
Login: timeoutConfig,
Registration: timeoutConfig,
},
})
}
/*
* 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 webauthn
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/x509"
"github.com/fxamacker/cbor/v2"
"github.com/go-webauthn/webauthn/protocol/webauthncose"
wan "github.com/go-webauthn/webauthn/webauthn"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
)
// curveP256CBOR is the constant for the P-256 curve in CBOR.
// https://datatracker.ietf.org/doc/html/rfc8152#section-13.1
const curveP256CBOR = 1
type credentialFlags struct {
BE, BS bool
}
func deviceToCredential(
dev *types.MFADevice,
idOnly bool,
currentFlags *credentialFlags,
) (wan.Credential, bool) {
switch dev := dev.Device.(type) {
case *types.MFADevice_U2F:
var pubKeyCBOR []byte
if !idOnly {
var err error
pubKeyCBOR, err = u2fDERKeyToCBOR(dev.U2F.PubKey)
if err != nil {
log.WarnContext(context.Background(), "failed to convert U2F device key to CBOR", "error", err)
return wan.Credential{}, false
}
}
return wan.Credential{
ID: dev.U2F.KeyHandle,
PublicKey: pubKeyCBOR,
Authenticator: wan.Authenticator{
SignCount: dev.U2F.Counter,
},
}, true
case *types.MFADevice_Webauthn:
var pubKeyCBOR []byte
if !idOnly {
pubKeyCBOR = dev.Webauthn.PublicKeyCbor
}
// Use BE/BS from the device, falling back to currentFlags for devices that
// haven't been backfilled yet.
var be, bs bool
if dev.Webauthn.CredentialBackupEligible != nil {
be = dev.Webauthn.CredentialBackupEligible.Value
} else {
be = currentFlags != nil && currentFlags.BE
}
if dev.Webauthn.CredentialBackedUp != nil {
bs = dev.Webauthn.CredentialBackedUp.Value
} else {
bs = currentFlags != nil && currentFlags.BS
}
return wan.Credential{
ID: dev.Webauthn.CredentialId,
PublicKey: pubKeyCBOR,
AttestationType: dev.Webauthn.AttestationType,
Flags: wan.CredentialFlags{
BackupEligible: be,
BackupState: bs,
},
Authenticator: wan.Authenticator{
AAGUID: dev.Webauthn.Aaguid,
SignCount: dev.Webauthn.SignatureCounter,
},
}, true
default:
return wan.Credential{}, false
}
}
func u2fDERKeyToCBOR(der []byte) ([]byte, error) {
pubKeyI, err := x509.ParsePKIXPublicKey(der)
if err != nil {
return nil, trace.Wrap(err)
}
// U2F device keys are guaranteed to be ECDSA/P256
// https://fidoalliance.org/specs/fido-u2f-v1.2-ps-20170411/fido-u2f-raw-message-formats-v1.2-ps-20170411.html#h3_registration-response-message-success.
pubKey, ok := pubKeyI.(*ecdsa.PublicKey)
if !ok {
return nil, trace.BadParameter("U2F public key has an unexpected type: %T", pubKeyI)
}
return U2FKeyToCBOR(pubKey)
}
// U2FKeyToCBOR transforms a DER-encoded U2F into its CBOR counterpart.
func U2FKeyToCBOR(pubKey *ecdsa.PublicKey) ([]byte, error) {
if pubKey.Curve != elliptic.P256() {
return nil, trace.BadParameter("unsupported curve %T", pubKey.Curve)
}
pubKeyBytes, err := pubKey.Bytes()
if err != nil {
return nil, trace.Wrap(err)
}
// First byte is the 0x04 prefix, which can be skipped. The rest are the x and y coordinates.
x, y := pubKeyBytes[1:33], pubKeyBytes[33:]
pubKeyCBOR, err := cbor.Marshal(&webauthncose.EC2PublicKeyData{
PublicKeyData: webauthncose.PublicKeyData{
KeyType: int64(webauthncose.EllipticKey),
Algorithm: int64(webauthncose.AlgES256),
},
Curve: curveP256CBOR,
XCoord: x,
YCoord: y,
})
return pubKeyCBOR, trace.Wrap(err)
}
/*
* 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 webauthn
import (
"bytes"
"context"
"encoding/base64"
"encoding/json"
"slices"
"sort"
"time"
"github.com/go-webauthn/webauthn/protocol"
wan "github.com/go-webauthn/webauthn/webauthn"
gogotypes "github.com/gogo/protobuf/types"
"github.com/gravitational/trace"
mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/auth/mfatypes"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
)
// loginIdentity contains the subset of services.Identity methods used by
// loginFlow.
type loginIdentity interface {
GetWebauthnLocalAuth(ctx context.Context, user string) (*types.WebauthnLocalAuth, error) // MFA
GetTeleportUserByWebauthnID(ctx context.Context, webID []byte) (string, error) // Passwordless
GetMFADevices(ctx context.Context, user string, withSecrets bool) ([]*types.MFADevice, error)
UpsertMFADevice(ctx context.Context, user string, d *types.MFADevice) error
}
// sessionIdentity abstracts operations over SessionData storage.
// * MFA uses per-user variants
// (services.Identity.Update/Get/DeleteWebauthnSessionData methods).
// * Passwordless uses global variants
// (services.Identity.Update/Get/DeleteGlobalWebauthnSessionData methods).
type sessionIdentity interface {
Upsert(ctx context.Context, user string, sd *wantypes.SessionData) error
Get(ctx context.Context, user string, challenge string) (*wantypes.SessionData, error)
Delete(ctx context.Context, user string, challenge string) error
}
// loginFlow implements both MFA and Passwordless authentication, exposing an
// interface that is the union of both login methods.
type loginFlow struct {
U2F *types.U2F
Webauthn *types.Webauthn
identity loginIdentity
sessionData sessionIdentity
}
func isReuseAllowedForScope(scope mfav1.ChallengeScope) bool {
switch scope {
case mfav1.ChallengeScope_CHALLENGE_SCOPE_ADMIN_ACTION,
mfav1.ChallengeScope_CHALLENGE_SCOPE_USER_SESSION,
mfav1.ChallengeScope_CHALLENGE_SCOPE_KUBE_LOCAL_PROXY_MULTI:
return true
default:
return false
}
}
func (f *loginFlow) begin(ctx context.Context, params BeginParams) (*wantypes.CredentialAssertion, error) {
if params.ChallengeExtensions == nil {
return nil, trace.BadParameter("requested challenge extensions must be supplied.")
}
if params.ChallengeExtensions.AllowReuse == mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_YES && !isReuseAllowedForScope(params.ChallengeExtensions.Scope) {
return nil, trace.BadParameter("mfa challenges with scope %s cannot allow reuse", params.ChallengeExtensions.Scope)
}
// discoverableLogin identifies logins started with an unknown/empty user.
discoverableLogin := params.ChallengeExtensions.Scope == mfav1.ChallengeScope_CHALLENGE_SCOPE_PASSWORDLESS_LOGIN
if params.User == "" && !discoverableLogin {
return nil, trace.BadParameter("user required")
}
var u *webUser
if discoverableLogin {
u = &webUser{} // Issue anonymous challenge.
} else {
webID, err := f.getWebID(ctx, params.User)
if err != nil {
return nil, trace.Wrap(err)
}
// Use existing devices to set the allowed credentials.
devices, err := f.identity.GetMFADevices(ctx, params.User, false /* withSecrets */)
if err != nil {
return nil, trace.Wrap(err)
}
// Filter devices with the wrong RPID and log an error.
foundInvalid := false
for i := 0; i < len(devices); i++ {
webDev := devices[i].GetWebauthn()
if webDev == nil || webDev.CredentialRpId == "" || webDev.CredentialRpId == f.Webauthn.RPID {
continue
}
const msg = "User device has unexpected RPID, excluding from allowed credentials. " +
"RPID changes are not supported by WebAuthn, this is likely to cause permanent authentication problems for your users. " +
"Consider reverting the change or reset your users so they may register their devices again."
log.ErrorContext(ctx, msg,
"user", params.User,
"device", devices[i].GetName(),
"rpid", webDev.CredentialRpId,
)
// "Cut" device from slice.
devices = slices.Delete(devices, i, i+1)
i--
foundInvalid = true
}
// Sort non-resident keys first, which may cause clients to favor them for
// MFA in some scenarios (eg, tsh).
sort.Slice(devices, func(i, j int) bool {
dev1, dev2 := devices[i], devices[j]
web1, web2 := dev1.GetWebauthn(), dev2.GetWebauthn()
resident1 := web1 != nil && web1.ResidentKey
resident2 := web2 != nil && web2.ResidentKey
return !resident1 && resident2
})
u = newWebUser(webUserOpts{
name: params.User,
webID: webID,
devices: devices,
credentialIDOnly: true,
})
// Let's make sure we have at least one registered credential here, since we
// have to allow zero credentials for passwordless below.
if len(u.credentials) == 0 {
if foundInvalid {
return nil, trace.Wrap(ErrInvalidCredentials)
}
return nil, trace.NotFound("found no credentials for user %q", params.User)
}
}
// TODO(codingllama): Use the "official" appid impl by duo-labs/webauthn.
var opts []wan.LoginOption
if f.U2F != nil && f.U2F.AppID != "" {
// See https://www.w3.org/TR/webauthn-2/#sctn-appid-extension.
opts = append(opts, wan.WithAssertionExtensions(protocol.AuthenticationExtensions{
wantypes.AppIDExtension: f.U2F.AppID,
}))
}
// Set the user verification requirement, if present, only for
// non-discoverable logins.
// For discoverable logins we rely on the wan.WebAuthn default set below.
if !discoverableLogin && params.ChallengeExtensions.UserVerificationRequirement != "" {
uvr := protocol.UserVerificationRequirement(params.ChallengeExtensions.UserVerificationRequirement)
opts = append(opts, wan.WithUserVerification(uvr))
}
// Create the WebAuthn object and issue a new challenge.
web, err := newWebAuthn(webAuthnParams{
cfg: f.Webauthn,
rpID: f.Webauthn.RPID,
requireUserVerification: discoverableLogin,
})
if err != nil {
return nil, trace.Wrap(err)
}
var assertion *protocol.CredentialAssertion
var sessionData *wan.SessionData
if discoverableLogin {
assertion, sessionData, err = web.BeginDiscoverableLogin(opts...)
} else {
assertion, sessionData, err = web.BeginLogin(u, opts...)
}
if err != nil {
return nil, trace.Wrap(err)
}
// Store SessionData - it's checked against the user response by Finish.
sd, err := wantypes.SessionDataFromProtocol(sessionData)
if err != nil {
return nil, trace.Wrap(err)
}
sd.ChallengeExtensions = &mfatypes.ChallengeExtensions{
Scope: params.ChallengeExtensions.Scope,
AllowReuse: params.ChallengeExtensions.AllowReuse,
UserVerificationRequirement: params.ChallengeExtensions.UserVerificationRequirement,
}
// Attach SIP if provided.
if params.SessionIdentifyingPayload != nil {
sd.Payload = &mfatypes.SessionIdentifyingPayload{
SSHSessionID: params.SessionIdentifyingPayload.GetSshSessionId(),
TLSSessionID: params.SessionIdentifyingPayload.GetTlsSessionId(),
}
}
// Attach source and target cluster names.
sd.SourceCluster, sd.TargetCluster = params.SourceCluster, params.TargetCluster
if err := f.sessionData.Upsert(ctx, params.User, sd); err != nil {
return nil, trace.Wrap(err)
}
return wantypes.CredentialAssertionFromProtocol(assertion), nil
}
func (f *loginFlow) getWebID(ctx context.Context, user string) ([]byte, error) {
wla, err := f.identity.GetWebauthnLocalAuth(ctx, user)
switch {
case trace.IsNotFound(err):
return nil, nil // OK, legacy U2F users may not have a webID.
case err != nil:
return nil, trace.Wrap(err)
}
return wla.UserID, nil
}
// LoginData is data gathered from a successful webauthn login.
type LoginData struct {
// User is the Teleport user.
User string
// Device is the MFA device used to authenticate the user.
Device *types.MFADevice
// AllowReuse is whether the webauthn challenge used for this login
// can be reused by the user for subsequent logins, until it expires.
AllowReuse mfav1.ChallengeAllowReuse
// Payload is the optional session identifying payload to attach to the login.
Payload *mfatypes.SessionIdentifyingPayload
// SourceCluster is the source cluster name associated with this login.
SourceCluster string
// TargetCluster is the target cluster name associated with this login.
TargetCluster string
}
func (f *loginFlow) finish(ctx context.Context, user string, resp *wantypes.CredentialAssertionResponse, requiredExtensions *mfav1.ChallengeExtensions) (*LoginData, error) {
if requiredExtensions == nil {
return nil, trace.BadParameter("requested challenge extensions must be supplied.")
}
discoverableLogin := requiredExtensions.Scope == mfav1.ChallengeScope_CHALLENGE_SCOPE_PASSWORDLESS_LOGIN
switch {
case user == "" && !discoverableLogin:
return nil, trace.BadParameter("user required")
case resp == nil:
// resp != nil is good enough to proceed, we leave remaining validations to
// duo-labs/webauthn.
return nil, trace.BadParameter("credential assertion response required")
}
parsedResp, err := parseCredentialResponse(resp)
if err != nil {
return nil, trace.Wrap(err)
}
origin := parsedResp.Response.CollectedClientData.Origin
if err := validateOrigin(origin, f.Webauthn.RPID); err != nil {
log.DebugContext(ctx, "origin validation failed", "error", err)
return nil, trace.Wrap(err)
}
var webID []byte
if discoverableLogin {
webID = parsedResp.Response.UserHandle
if len(webID) == 0 {
return nil, trace.BadParameter("webauthn user handle required for passwordless")
}
// Fetch user from WebAuthn UserHandle (aka User ID).
teleportUser, err := f.identity.GetTeleportUserByWebauthnID(ctx, webID)
if err != nil {
return nil, trace.Wrap(err)
}
user = teleportUser
} else {
webID, err = f.getWebID(ctx, user)
if err != nil {
return nil, trace.Wrap(err)
}
}
// Find the device used to sign the credentials. It must be a previously
// registered device.
devices, err := f.identity.GetMFADevices(ctx, user, false /* withSecrets */)
if err != nil {
return nil, trace.Wrap(err)
}
dev, ok := findDeviceByID(devices, parsedResp.RawID)
if !ok {
return nil, trace.BadParameter(
"unknown device credential: %q", base64.RawURLEncoding.EncodeToString(parsedResp.RawID))
}
// Is an U2F device trying to login? If yes, use RPID = App ID.
// Technically browsers should reply with the appid extension set to true[1],
// but in actuality they don't send anything.
// [1] https://www.w3.org/TR/webauthn-2/#sctn-appid-extension.
rpID := f.Webauthn.RPID
switch {
case dev.GetU2F() != nil && f.U2F == nil:
return nil, trace.BadParameter("U2F device attempted login, but U2F configuration not present")
case dev.GetU2F() != nil:
rpID = f.U2F.AppID
}
u := newWebUser(webUserOpts{
name: user,
webID: webID,
devices: []*types.MFADevice{dev},
currentFlags: &credentialFlags{
BE: parsedResp.Response.AuthenticatorData.Flags.HasBackupEligible(),
BS: parsedResp.Response.AuthenticatorData.Flags.HasBackupState(),
},
})
// Fetch the previously-stored SessionData, so it's checked against the user
// response.
challenge := parsedResp.Response.CollectedClientData.Challenge
sd, err := f.sessionData.Get(ctx, user, challenge)
if err != nil {
return nil, trace.Wrap(err)
}
// Check if the given scope is satisfied by the challenge scope.
if requiredExtensions.Scope != sd.ChallengeExtensions.Scope && requiredExtensions.Scope != mfav1.ChallengeScope_CHALLENGE_SCOPE_UNSPECIFIED {
return nil, trace.AccessDenied("required scope %q is not satisfied by the given webauthn session with scope %q", requiredExtensions.Scope, sd.ChallengeExtensions.Scope)
}
noReuseAllowed := requiredExtensions.AllowReuse == mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_NO
challengeAllowReuse := sd.ChallengeExtensions.AllowReuse == mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_YES
// If this session is reusable, but this login forbids reusable sessions, return an error.
if noReuseAllowed && challengeAllowReuse {
return nil, trace.AccessDenied("the given webauthn session allows reuse, but reuse is not permitted in this context")
}
// Verify (and possibly correct) the user verification requirement.
// A mismatch here could indicate a programming error or even foul play.
uvr := protocol.UserVerificationRequirement(requiredExtensions.UserVerificationRequirement)
if (discoverableLogin || uvr == protocol.VerificationRequired) && sd.UserVerification != string(protocol.VerificationRequired) {
// This is not a failure yet, but will likely become one.
sd.UserVerification = string(protocol.VerificationRequired)
const msg = "User verification required by extensions but not by challenge. " +
"Increased SessionData.UserVerification."
log.WarnContext(ctx, msg, "user_verification", sd.UserVerification)
}
sessionData := wantypes.SessionDataToProtocol(sd)
// Make sure _all_ credentials in the session are accounted for by the user.
// webauthn.ValidateLogin requires it.
for _, allowedCred := range sessionData.AllowedCredentialIDs {
if bytes.Equal(parsedResp.RawID, allowedCred) {
continue
}
u.credentials = append(u.credentials, wan.Credential{
ID: allowedCred,
})
}
// Create a WebAuthn matching the expected RPID and Origin, then verify the
// signed challenge.
web, err := newWebAuthn(webAuthnParams{
cfg: f.Webauthn,
rpID: rpID,
origin: origin,
requireUserVerification: discoverableLogin,
})
if err != nil {
return nil, trace.Wrap(err)
}
var credential *wan.Credential
if discoverableLogin {
discoverUser := func(_, _ []byte) (wan.User, error) { return u, nil }
credential, err = web.ValidateDiscoverableLogin(discoverUser, *sessionData, parsedResp)
} else {
credential, err = web.ValidateLogin(u, *sessionData, parsedResp)
}
if err != nil {
return nil, trace.Wrap(err)
}
if credential.Authenticator.CloneWarning {
// Reused challenges trigger clone warnings because, after first use, we expect
// the counter to be up by one.
isReuseCounterMismatch := challengeAllowReuse &&
len(u.credentials) == 1 /* sanity check */ &&
credential.Authenticator.SignCount != u.credentials[0].Authenticator.SignCount-1
if !isReuseCounterMismatch {
log.WarnContext(ctx, "Clone warning detected for device, the device counter may be malfunctioning",
"user", user,
"device", dev.GetName(),
)
}
}
// Update last used timestamp and device counter.
if err := updateCredentialAndTimestamps(dev, credential, discoverableLogin); err != nil {
return nil, trace.Wrap(err)
}
// Retroactively write the credential RPID, now that it cleared authn.
if webDev := dev.GetWebauthn(); webDev != nil && webDev.CredentialRpId == "" {
log.DebugContext(ctx, "Recording RPID in device",
"rpid", rpID,
"user", user,
"device", dev.GetName(),
)
webDev.CredentialRpId = rpID
}
if err := f.identity.UpsertMFADevice(ctx, user, dev); err != nil {
return nil, trace.Wrap(err)
}
// The user just solved the challenge, so let's make sure it won't be used
// again, unless reuse is explicitly allowed.
// Note that even reusable sessions are deleted when their expiration time
// passes.
if !challengeAllowReuse {
if err := f.sessionData.Delete(ctx, user, challenge); err != nil {
log.WarnContext(ctx, "failed to delete login SessionData for user",
"user", user,
"scope", sd.ChallengeExtensions.Scope,
)
}
}
return &LoginData{
User: user,
Device: dev,
AllowReuse: sd.ChallengeExtensions.AllowReuse,
Payload: sd.Payload,
SourceCluster: sd.SourceCluster,
TargetCluster: sd.TargetCluster,
}, nil
}
func parseCredentialResponse(resp *wantypes.CredentialAssertionResponse) (*protocol.ParsedCredentialAssertionData, error) {
// Do not pass extensions on to duo-labs/webauthn, they won't go past JSON
// unmarshal.
exts := resp.Extensions
resp.Extensions = nil
defer func() { resp.Extensions = exts }()
// This is a roundabout way of getting resp validated, but unfortunately the
// APIs don't provide a better method (and it seems better than duplicating
// library code here).
body, err := json.Marshal(resp)
if err != nil {
return nil, trace.Wrap(err)
}
return protocol.ParseCredentialRequestResponseBody(bytes.NewReader(body))
}
func findDeviceByID(devices []*types.MFADevice, id []byte) (*types.MFADevice, bool) {
for _, dev := range devices {
switch d := dev.Device.(type) {
case *types.MFADevice_U2F:
if bytes.Equal(d.U2F.KeyHandle, id) {
return dev, true
}
case *types.MFADevice_Webauthn:
if bytes.Equal(d.Webauthn.CredentialId, id) {
return dev, true
}
}
}
return nil, false
}
func updateCredentialAndTimestamps(
dest *types.MFADevice,
credential *wan.Credential,
discoverableLogin bool,
) error {
switch d := dest.Device.(type) {
case *types.MFADevice_U2F:
d.U2F.Counter = credential.Authenticator.SignCount
case *types.MFADevice_Webauthn:
d.Webauthn.SignatureCounter = credential.Authenticator.SignCount
// Backfill ResidentKey field.
// This may happen if an authenticator created for "MFA" was actually
// resident all along (eg, Safari/Touch ID registrations).
if discoverableLogin && !d.Webauthn.ResidentKey {
d.Webauthn.ResidentKey = true
}
// Backfill BE/BS bits.
if d.Webauthn.CredentialBackupEligible == nil {
d.Webauthn.CredentialBackupEligible = &gogotypes.BoolValue{
Value: credential.Flags.BackupEligible,
}
log.DebugContext(context.Background(), "Backfilled Webauthn device BE flag",
"device", dest.GetName(),
"be", credential.Flags.BackupEligible,
)
}
if d.Webauthn.CredentialBackedUp == nil {
d.Webauthn.CredentialBackedUp = &gogotypes.BoolValue{
Value: credential.Flags.BackupState,
}
log.DebugContext(context.Background(), "Backfilled Webauthn device BS flag",
"device", dest.GetName(),
"bs", credential.Flags.BackupState,
)
}
default:
return trace.BadParameter("unexpected device type for webauthn: %T", d)
}
dest.LastUsed = time.Now()
return nil
}
/*
* 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 webauthn
import (
"context"
"errors"
"github.com/gravitational/trace"
mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1"
mfav2 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v2"
"github.com/gravitational/teleport/api/types"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
)
// ErrInvalidCredentials is a special kind of credential "NotFound" error, where
// the user has only devices registered to other RPIDs.
// Possible fixes include reseting the affected users (likely the entire
// cluster), or rolling back to a good WebAuthn configuration (if still
// possible).
var ErrInvalidCredentials = errors.New("user has only invalid WebAuthn registrations, consider a user reset")
// LoginIdentity represents the subset of Identity methods used by LoginFlow.
// It exists to better scope LoginFlow's use of Identity and to facilitate
// testing.
type LoginIdentity interface {
GetWebauthnLocalAuth(ctx context.Context, user string) (*types.WebauthnLocalAuth, error)
GetMFADevices(ctx context.Context, user string, withSecrets bool) ([]*types.MFADevice, error)
UpsertMFADevice(ctx context.Context, user string, d *types.MFADevice) error
UpsertWebauthnSessionData(ctx context.Context, user, sessionID string, sd *wantypes.SessionData) error
GetWebauthnSessionData(ctx context.Context, user, sessionID string) (*wantypes.SessionData, error)
DeleteWebauthnSessionData(ctx context.Context, user, sessionID string) error
}
// WithDevices returns a LoginIdentity backed by a fixed set of devices.
// The supplied devices are returned in all GetMFADevices calls.
func WithDevices(identity LoginIdentity, devs []*types.MFADevice) LoginIdentity {
return &loginWithDevices{
LoginIdentity: identity,
devices: devs,
}
}
type loginWithDevices struct {
LoginIdentity
devices []*types.MFADevice
}
func (l *loginWithDevices) GetMFADevices(_ context.Context, _ string, _ bool) ([]*types.MFADevice, error) {
return l.devices, nil
}
// LoginFlow represents the WebAuthn login procedure (aka authentication).
//
// The login flow consists of:
//
// 1. Client requests a CredentialAssertion (containing, among other info, a
// challenge to be signed)
// 2. Server runs Begin(), generates a credential assertion.
// 3. Client validates the assertion, performs a user presence test (usually by
// asking the user to touch a secure token), and replies with
// CredentialAssertionResponse (containing the signed challenge)
// 4. Server runs Finish()
// 5. If all server-side checks are successful, then login/authentication is
// complete.
//
// LoginFlow is used in the following scenarios:
// - Password plus challenge logins
// - Presence verification checks (eg, session MFA)
// - User verification checks after the initial login (eg, password changes
// with only a discoverable credential).
type LoginFlow struct {
U2F *types.U2F
Webauthn *types.Webauthn
// Identity is typically an implementation of the Identity service, ie, an
// object with access to user, device and MFA storage.
Identity LoginIdentity
}
// BeginParams contains parameters for the Begin method.
type BeginParams struct {
User string
ChallengeExtensions *mfav1.ChallengeExtensions
SessionIdentifyingPayload *mfav2.SessionIdentifyingPayload
SourceCluster string
TargetCluster string
}
// Begin is the first step of the LoginFlow.
// The CredentialAssertion created is relayed back to the client, who in turn
// performs a user presence check and signs the challenge contained within the
// assertion.
// As a side effect Begin may assign (and record in storage) a WebAuthn ID for
// the user.
// Requested challenge extensions will be stored on the stored webauthn challenge
// record. These extensions indicate additional rules/properties of the webauthn
// challenge that can be validated in the final login step.
func (f *LoginFlow) Begin(ctx context.Context, params BeginParams) (*wantypes.CredentialAssertion, error) {
// Disallow passwordless through here.
// lf.begin() does other challengeExtensions checks, including `nil`.
if params.ChallengeExtensions != nil && params.ChallengeExtensions.Scope == mfav1.ChallengeScope_CHALLENGE_SCOPE_PASSWORDLESS_LOGIN {
return nil, trace.BadParameter("passwordless challenge scope is not allowed for MFA flows")
}
lf := &loginFlow{
U2F: f.U2F,
Webauthn: f.Webauthn,
identity: mfaIdentity{f.Identity},
// TODO(codingllama): Record session data to distinct scope keys based on
// the actual challenge scope.
sessionData: (*userSessionStorage)(f),
}
return lf.begin(ctx, params)
}
// Finish is the second and last step of the LoginFlow.
// Expected challenge extensions will be validated against the stored webauthn
// challenge record.
// It returns the MFADevice used to solve the challenge, the associated Teleport
// user name, and other login properties. If login is successful, Finish has the
// side effect of updating the counter and last used timestamp of the MFADevice
// used.
func (f *LoginFlow) Finish(ctx context.Context, user string, resp *wantypes.CredentialAssertionResponse, requiredExtensions *mfav1.ChallengeExtensions) (*LoginData, error) {
lf := &loginFlow{
U2F: f.U2F,
Webauthn: f.Webauthn,
identity: mfaIdentity{f.Identity},
sessionData: (*userSessionStorage)(f),
}
return lf.finish(ctx, user, resp, requiredExtensions)
}
type mfaIdentity struct {
LoginIdentity
}
func (m mfaIdentity) GetTeleportUserByWebauthnID(_ context.Context, _ []byte) (string, error) {
return "", errors.New("lookup by webauthn ID not supported for MFA")
}
// userSessionStorage implements sessionIdentity using LoginFlow.
type userSessionStorage LoginFlow
func (s *userSessionStorage) Upsert(ctx context.Context, user string, sd *wantypes.SessionData) error {
return s.Identity.UpsertWebauthnSessionData(ctx, user, scopeLogin, sd)
}
func (s *userSessionStorage) Get(ctx context.Context, user string, _ string) (*wantypes.SessionData, error) {
return s.Identity.GetWebauthnSessionData(ctx, user, scopeLogin)
}
func (s *userSessionStorage) Delete(ctx context.Context, user string, _ string) error {
return s.Identity.DeleteWebauthnSessionData(ctx, user, scopeLogin)
}
/*
* 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 webauthn
import (
"context"
"encoding/base64"
"errors"
mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1"
"github.com/gravitational/teleport/api/types"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
)
// PasswordlessIdentity represents the subset of Identity methods used by
// PasswordlessFlow.
type PasswordlessIdentity interface {
GetMFADevices(ctx context.Context, user string, withSecrets bool) ([]*types.MFADevice, error)
UpsertMFADevice(ctx context.Context, user string, d *types.MFADevice) error
UpsertGlobalWebauthnSessionData(ctx context.Context, scope, id string, sd *wantypes.SessionData) error
GetGlobalWebauthnSessionData(ctx context.Context, scope, id string) (*wantypes.SessionData, error)
DeleteGlobalWebauthnSessionData(ctx context.Context, scope, id string) error
GetTeleportUserByWebauthnID(ctx context.Context, webID []byte) (string, error)
}
// PasswordlessFlow provides passwordless authentication.
//
// PasswordlessFlow is used mainly for the initial passwordless login.
// For UV=1 assertions after login, use [LoginFlow.Begin] with the desired
// [mfav1.ChallengeExtensions.UserVerificationRequirement].
type PasswordlessFlow struct {
Webauthn *types.Webauthn
Identity PasswordlessIdentity
}
// Begin is the first step of the passwordless login flow.
// It works similarly to LoginFlow.Begin, but it doesn't require a Teleport
// username nor implies a previous password-validation step.
func (f *PasswordlessFlow) Begin(ctx context.Context) (*wantypes.CredentialAssertion, error) {
lf := &loginFlow{
Webauthn: f.Webauthn,
identity: passwordlessIdentity{f.Identity},
sessionData: (*globalSessionStorage)(f),
}
chalExt := &mfav1.ChallengeExtensions{
Scope: mfav1.ChallengeScope_CHALLENGE_SCOPE_PASSWORDLESS_LOGIN,
AllowReuse: mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_NO,
}
return lf.begin(ctx, BeginParams{
User: "",
ChallengeExtensions: chalExt,
})
}
// Finish is the last step of the passwordless login flow.
// It works similarly to LoginFlow.Finish, but the user identity is established
// via the response UserHandle, instead of an explicit Teleport username.
func (f *PasswordlessFlow) Finish(ctx context.Context, resp *wantypes.CredentialAssertionResponse) (*LoginData, error) {
lf := &loginFlow{
Webauthn: f.Webauthn,
identity: passwordlessIdentity{f.Identity},
sessionData: (*globalSessionStorage)(f),
}
requiredExt := &mfav1.ChallengeExtensions{
Scope: mfav1.ChallengeScope_CHALLENGE_SCOPE_PASSWORDLESS_LOGIN,
AllowReuse: mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_NO,
}
return lf.finish(ctx, "" /* user */, resp, requiredExt)
}
type passwordlessIdentity struct {
PasswordlessIdentity
}
func (p passwordlessIdentity) UpsertWebauthnLocalAuth(ctx context.Context, user string, wla *types.WebauthnLocalAuth) error {
return errors.New("webauthn local auth not supported for passwordless")
}
func (p passwordlessIdentity) GetWebauthnLocalAuth(ctx context.Context, user string) (*types.WebauthnLocalAuth, error) {
return nil, errors.New("webauthn local auth not supported for passwordless")
}
type globalSessionStorage PasswordlessFlow
func (g *globalSessionStorage) Upsert(ctx context.Context, user string, sd *wantypes.SessionData) error {
id := base64.RawURLEncoding.EncodeToString(sd.Challenge)
return g.Identity.UpsertGlobalWebauthnSessionData(ctx, scopeLogin, id, sd)
}
func (g *globalSessionStorage) Get(ctx context.Context, user string, challenge string) (*wantypes.SessionData, error) {
return g.Identity.GetGlobalWebauthnSessionData(ctx, scopeLogin, challenge)
}
func (g *globalSessionStorage) Delete(ctx context.Context, user string, challenge string) error {
return g.Identity.DeleteGlobalWebauthnSessionData(ctx, scopeLogin, challenge)
}
/*
* 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 webauthn
import (
"net/url"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/utils"
)
func validateOrigin(origin, rpID string) error {
parsedOrigin, err := url.Parse(origin)
if err != nil {
return trace.BadParameter("origin is not a valid URL: %v", err)
}
host, err := utils.Host(parsedOrigin.Host)
if err != nil {
return trace.BadParameter("extracting host from origin: %v", err)
}
// TODO(codingllama): Check origin against the public addresses of Proxies and
// Auth Servers
// Accept origins whose host matches the RPID.
if host == rpID {
return nil
}
// Accept origins whose host is a subdomain of RPID.
originParts := strings.Split(host, ".")
rpParts := strings.Split(rpID, ".")
if len(originParts) <= len(rpParts) {
return trace.BadParameter("origin doesn't match RPID")
}
i := len(originParts) - 1
j := len(rpParts) - 1
for j >= 0 {
if originParts[i] != rpParts[j] {
return trace.BadParameter("origin doesn't match RPID")
}
i--
j--
}
return nil
}
/*
* 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 webauthn
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"sync"
"time"
"github.com/go-webauthn/webauthn/protocol"
wan "github.com/go-webauthn/webauthn/webauthn"
gogotypes "github.com/gogo/protobuf/types"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
)
// RegistrationIdentity represents the subset of Identity methods used by
// RegistrationFlow.
type RegistrationIdentity interface {
UpsertWebauthnLocalAuth(ctx context.Context, user string, wla *types.WebauthnLocalAuth) error
GetWebauthnLocalAuth(ctx context.Context, user string) (*types.WebauthnLocalAuth, error)
GetTeleportUserByWebauthnID(ctx context.Context, webID []byte) (string, error)
GetMFADevices(ctx context.Context, user string, withSecrets bool) ([]*types.MFADevice, error)
UpsertMFADevice(ctx context.Context, user string, d *types.MFADevice) error
UpsertWebauthnSessionData(ctx context.Context, user, sessionID string, sd *wantypes.SessionData) error
GetWebauthnSessionData(ctx context.Context, user, sessionID string) (*wantypes.SessionData, error)
DeleteWebauthnSessionData(ctx context.Context, user, sessionID string) error
}
// WithInMemorySessionData returns a RegistrationIdentity implementation that
// keeps SessionData in memory.
func WithInMemorySessionData(identity RegistrationIdentity) RegistrationIdentity {
return &inMemoryIdentity{
RegistrationIdentity: identity,
sessionData: make(map[string]*wantypes.SessionData),
}
}
type inMemoryIdentity struct {
RegistrationIdentity
// mu guards the fields below it.
// We don't foresee concurrent use for inMemoryIdentity, but it's easy enough
// to play it safe.
mu sync.RWMutex
sessionData map[string]*wantypes.SessionData
}
func (identity *inMemoryIdentity) UpsertWebauthnSessionData(ctx context.Context, user, sessionID string, sd *wantypes.SessionData) error {
identity.mu.Lock()
defer identity.mu.Unlock()
identity.sessionData[sessionDataKey(user, sessionID)] = sd
return nil
}
func (identity *inMemoryIdentity) GetWebauthnSessionData(ctx context.Context, user, sessionID string) (*wantypes.SessionData, error) {
identity.mu.RLock()
defer identity.mu.RUnlock()
sd, ok := identity.sessionData[sessionDataKey(user, sessionID)]
if !ok {
return nil, trace.NotFound("session data for user %v not found ", user)
}
// The only known caller of GetWebauthnSessionData is the webauthn package
// itself, so we trust it to not modify the SessionData we are handing back.
return sd, nil
}
func (identity *inMemoryIdentity) DeleteWebauthnSessionData(ctx context.Context, user, sessionID string) error {
key := sessionDataKey(user, sessionID)
identity.mu.Lock()
defer identity.mu.Unlock()
if _, ok := identity.sessionData[key]; !ok {
return trace.NotFound("session data for user %v not found ", user)
}
delete(identity.sessionData, key)
return nil
}
func sessionDataKey(user, sessionID string) string {
return fmt.Sprintf("%v/%v", user, sessionID)
}
// RegistrationFlow represents the WebAuthn registration ceremony.
//
// Registration consists of:
//
// 1. Client requests a CredentialCreation (containing a challenge and various
// settings that may constrain allowed authenticators).
// 2. Server runs Begin(), generates a credential creation.
// 3. Client validates the credential creation, performs a user presence test
// (usually by asking the user to touch a secure token), and replies with a
// CredentialCreationResponse (containing the signed challenge and
// information about the credential and authenticator)
// 4. Server runs Finish()
// 5. If all server-side checks are successful, then registration is complete
// and the authenticator may now be used to login.
type RegistrationFlow struct {
Webauthn *types.Webauthn
Identity RegistrationIdentity
}
// Begin is the first step of the registration ceremony.
// The CredentialCreation created is relayed back to the client, who in turn
// performs a user presence check and signs the challenge contained within it.
// If passwordless is set, then registration asks the authenticator for a
// resident key.
// As a side effect Begin may assign (and record in storage) a WebAuthn ID for
// the user.
func (f *RegistrationFlow) Begin(ctx context.Context, user string, passwordless bool) (*wantypes.CredentialCreation, error) {
if user == "" {
return nil, trace.BadParameter("user required")
}
// Exclude known devices from the ceremony.
devices, err := f.Identity.GetMFADevices(ctx, user, false /* withSecrets */)
if err != nil {
return nil, trace.Wrap(err)
}
var exclusions []protocol.CredentialDescriptor
for _, dev := range devices {
// Skip existing U2F devices, letting users "upgrade" their registration is
// good for us.
if dev.GetU2F() != nil {
continue
}
// Let authenticator "upgrades" from non-resident (MFA) to resident
// (passwordless) happen, but prevent "downgrades" from resident to
// non-resident.
//
// Modern passkey implementations will "disobey" our MFA registrations and
// actually create passkeys, silently replacing the old passkey with the new
// "MFA" key, which can make Teleport confused (for example, by letting the
// "MFA" key be deleted because Teleport thinks the passkey still exists).
if webDev := dev.GetWebauthn(); webDev != nil && !webDev.ResidentKey && passwordless {
continue
}
cred, ok := deviceToCredential(dev, true /* idOnly */, nil /* currentFlags */)
if !ok {
continue
}
exclusions = append(exclusions, protocol.CredentialDescriptor{
Type: protocol.PublicKeyCredentialType,
CredentialID: cred.ID,
})
}
webID, err := upsertOrGetWebID(ctx, user, f.Identity)
if err != nil {
return nil, trace.Wrap(err)
}
u := newWebUser(webUserOpts{
name: user,
webID: webID,
credentialIDOnly: true,
})
web, err := newWebAuthn(webAuthnParams{
cfg: f.Webauthn,
rpID: f.Webauthn.RPID,
requireResidentKey: passwordless,
requireUserVerification: passwordless,
})
if err != nil {
return nil, trace.Wrap(err)
}
cc, sessionData, err := web.BeginRegistration(
u,
wan.WithExclusions(exclusions),
wan.WithExtensions(protocol.AuthenticationExtensions{
// Query authenticator on whether the resulting credential is resident,
// despite our requirements.
wantypes.CredPropsExtension: true,
}),
)
if err != nil {
return nil, trace.Wrap(err)
}
// TODO(codingllama): Send U2F App ID back in creation requests too. Useful to
// detect duplicate devices.
sd, err := wantypes.SessionDataFromProtocol(sessionData)
if err != nil {
return nil, trace.Wrap(err)
}
if err := f.Identity.UpsertWebauthnSessionData(ctx, user, scopeSession, sd); err != nil {
return nil, trace.Wrap(err)
}
return wantypes.CredentialCreationFromProtocol(cc), nil
}
func upsertOrGetWebID(ctx context.Context, user string, identity RegistrationIdentity) ([]byte, error) {
wla, err := identity.GetWebauthnLocalAuth(ctx, user)
switch {
case trace.IsNotFound(err): // first-time user, create a new ID
webID := []byte(uuid.New().String())
err := identity.UpsertWebauthnLocalAuth(ctx, user, &types.WebauthnLocalAuth{
UserID: webID,
})
return webID[:], trace.Wrap(err)
case err != nil:
return nil, trace.Wrap(err)
}
// Attempt to fix the webID->user index, if necessary.
// This applies to legacy (Teleport 8.x) registrations and to eventual bad
// writes.
indexedUser, err := identity.GetTeleportUserByWebauthnID(ctx, wla.UserID)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
if indexedUser != user {
// Re-write wla to force an index update.
if err := identity.UpsertWebauthnLocalAuth(ctx, user, wla); err != nil {
return nil, trace.Wrap(err)
}
}
return wla.UserID, nil
}
// RegisterResponse represents fields needed to finish registering a new
// WebAuthn device.
type RegisterResponse struct {
// User is the device owner.
User string
// DeviceName is the name for the new device.
DeviceName string
// CreationResponse is the response from the new device.
CreationResponse *wantypes.CredentialCreationResponse
// Passwordless is true if this is expected to be a passwordless registration.
// Callers may make certain concessions when processing passwordless
// registration (such as skipping password validation), this flag reflects that.
// The data stored in the Begin SessionData must match the passwordless flag,
// otherwise the registration is denied.
Passwordless bool
}
// Finish is the second and last step of the registration ceremony.
// If successful, it returns the created MFADevice. Finish has the side effect
// or writing the device to storage (using its Identity interface).
func (f *RegistrationFlow) Finish(ctx context.Context, req RegisterResponse) (*types.MFADevice, error) {
switch {
case req.User == "":
return nil, trace.BadParameter("user required")
case req.DeviceName == "":
return nil, trace.BadParameter("device name required")
case req.CreationResponse == nil:
return nil, trace.BadParameter("credential creation response required")
}
parsedResp, err := parseCredentialCreationResponse(req.CreationResponse)
if err != nil {
return nil, trace.Wrap(err)
}
origin := parsedResp.Response.CollectedClientData.Origin
if err := validateOrigin(origin, f.Webauthn.RPID); err != nil {
log.DebugContext(ctx, "origin validation failed", "error", err)
return nil, trace.Wrap(err)
}
// TODO(codingllama): Verify that the public key matches the allowed
// credential params? It doesn't look like duo-labs/webauthn does that.
wla, err := f.Identity.GetWebauthnLocalAuth(ctx, req.User)
if err != nil {
return nil, trace.Wrap(err)
}
u := newWebUser(webUserOpts{
name: req.User,
webID: wla.UserID,
credentialIDOnly: true,
})
sd, err := f.Identity.GetWebauthnSessionData(ctx, req.User, scopeSession)
if err != nil {
return nil, trace.Wrap(err)
}
sessionData := wantypes.SessionDataToProtocol(sd)
// Activate passwordless switches (resident key, user verification) if we
// required verification in the begin step.
passwordless := sessionData.UserVerification == protocol.VerificationRequired
if req.Passwordless && !passwordless {
return nil, trace.BadParameter("passwordless registration failed, requested CredentialCreation was for an MFA registration")
}
web, err := newWebAuthn(webAuthnParams{
cfg: f.Webauthn,
rpID: f.Webauthn.RPID,
origin: origin,
requireResidentKey: passwordless,
requireUserVerification: passwordless,
})
if err != nil {
return nil, trace.Wrap(err)
}
credential, err := web.CreateCredential(u, *sessionData, parsedResp)
if err != nil {
// Use a more friendly message for certain verification errors.
protocolErr := &protocol.Error{}
if errors.As(err, &protocolErr) &&
protocolErr.Type == protocol.ErrVerification.Type &&
passwordless &&
!parsedResp.Response.AttestationObject.AuthData.Flags.UserVerified() {
log.DebugContext(ctx, "WebAuthn: Replacing verification error with PIN message", "error", err)
return nil, trace.BadParameter("authenticator doesn't support passwordless, setting up a PIN may fix this")
}
return nil, trace.Wrap(err)
}
// Finally, check against attestation settings, if any.
// This runs after web.CreateCredential so we can take advantage of the
// attestation format validation it performs.
if err := verifyAttestation(f.Webauthn, parsedResp.Response.AttestationObject); err != nil {
return nil, trace.Wrap(err)
}
newDevice, err := types.NewMFADevice(req.DeviceName, uuid.NewString() /* id */, time.Now() /* addedAt */, &types.MFADevice_Webauthn{
Webauthn: &types.WebauthnDevice{
CredentialId: credential.ID,
PublicKeyCbor: credential.PublicKey,
AttestationType: credential.AttestationType,
Aaguid: credential.Authenticator.AAGUID,
SignatureCounter: credential.Authenticator.SignCount,
AttestationObject: req.CreationResponse.AttestationResponse.AttestationObject,
ResidentKey: req.Passwordless || hasCredPropsRK(req.CreationResponse),
CredentialRpId: f.Webauthn.RPID,
CredentialBackupEligible: &gogotypes.BoolValue{
Value: credential.Flags.BackupEligible,
},
CredentialBackedUp: &gogotypes.BoolValue{
Value: credential.Flags.BackupState,
},
},
})
if err != nil {
return nil, trace.Wrap(err)
}
// We delegate a few checks to identity, including:
// * The validity of the created MFADevice
// * Uniqueness validation of the deviceName
// * Uniqueness validation of the Webauthn credential ID.
if err := f.Identity.UpsertMFADevice(ctx, req.User, newDevice); err != nil {
return nil, trace.Wrap(err)
}
// Registration complete, remove the registration challenge we just used.
if err := f.Identity.DeleteWebauthnSessionData(ctx, req.User, scopeSession); err != nil {
log.WarnContext(ctx, "failed to delete registration SessionData for user", "user", req.User, "error", err)
}
return newDevice, nil
}
func parseCredentialCreationResponse(resp *wantypes.CredentialCreationResponse) (*protocol.ParsedCredentialCreationData, error) {
// Remove extensions before marshaling, duo-labs/webauthn isn't expecting it.
exts := resp.Extensions
resp.Extensions = nil
defer func() {
resp.Extensions = exts
}()
// This is a roundabout way of getting resp validated, but unfortunately the
// APIs don't provide a better method (and it seems better than duplicating
// library code here).
body, err := json.Marshal(resp)
if err != nil {
return nil, trace.Wrap(err)
}
parsedResp, err := protocol.ParseCredentialCreationResponseBody(bytes.NewReader(body))
return parsedResp, trace.Wrap(err)
}
func hasCredPropsRK(ccr *wantypes.CredentialCreationResponse) bool {
return ccr != nil &&
ccr.Extensions != nil &&
ccr.Extensions.CredProps != nil &&
ccr.Extensions.CredProps.RK
}
/*
* 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 webauthn
import (
wan "github.com/go-webauthn/webauthn/webauthn"
"github.com/gravitational/teleport/api/types"
)
// webUser implements a WebAuthn protocol user.
// It is used to provide user information to WebAuthn APIs, but has no direct
// counterpart in storage nor in other packages.
type webUser struct {
credentials []wan.Credential
name string
webID []byte
}
type webUserOpts struct {
name string
webID []byte
devices []*types.MFADevice
credentialIDOnly bool
currentFlags *credentialFlags
}
func newWebUser(opts webUserOpts) *webUser {
var credentials []wan.Credential
for _, dev := range opts.devices {
c, ok := deviceToCredential(dev, opts.credentialIDOnly, opts.currentFlags)
if ok {
credentials = append(credentials, c)
}
}
return &webUser{
credentials: credentials,
name: opts.name,
webID: opts.webID,
}
}
func (w *webUser) WebAuthnID() []byte {
return w.webID
}
func (w *webUser) WebAuthnName() string {
return w.name
}
func (w *webUser) WebAuthnDisplayName() string {
return w.name
}
func (w *webUser) WebAuthnIcon() string {
return ""
}
func (w *webUser) WebAuthnCredentials() []wan.Credential {
return w.credentials
}
// Teleport
// Copyright (C) 2026 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 sse
import (
"bufio"
"bytes"
"errors"
"io"
"iter"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
)
var (
// ErrEventTooLarge is the error returned when SSE event is larger than
// [MaxReadEventSize].
ErrEventTooLarge = errors.New("sse event exceeded max size")
)
// MaxReadEventSize defines the max size that can be read from an [io.Reader]
// to complete the event.
const MaxReadEventSize = teleport.MaxHTTPResponseSize
// Event is a server-sent event.
type Event struct {
Event string
ID string
Data []byte
Retry string
}
// Empty reports whether the event is empty.
func (e Event) Empty() bool {
return e.len() == 0
}
// Equal returns `true` if both events have the same contents.
func (e Event) Equal(b Event) bool {
return e.Event == b.Event && e.ID == b.ID &&
bytes.Equal(e.Data, b.Data) && e.Retry == b.Retry
}
func (e Event) len() int {
return len(e.Event) + len(e.ID) + len(e.Data) + len(e.Retry)
}
// Validates the fields of the SSE event.
func (e Event) validate() error {
switch {
case strings.ContainsAny(e.Event, "\r\n"):
return trace.BadParameter("Event field cannot contain line feed or carriage return")
case strings.ContainsAny(e.ID, "\r\n"):
return trace.BadParameter("ID field cannot contain line feed or carriage return")
case strings.ContainsAny(e.Retry, "\r\n"):
return trace.BadParameter("Retry field cannot contain line feed or carriage return")
case !onlyDigits(e.Retry):
return trace.BadParameter("Retry field can only contain digits")
}
return nil
}
// ReadEvents reads SSE events from the provided reader.
func ReadEvents(r io.Reader) iter.Seq2[Event, error] {
return func(yield func(Event, error) bool) {
var currentEvent *Event
scanner := bufio.NewScanner(r)
scanner.Buffer(nil, MaxReadEventSize)
scanner.Split(scanLines())
yieldEvent := func() bool {
// This handles cases where the stream has empty lines, so we
// consumed them, and there is no need to invoke `yield`.
if currentEvent == nil {
return true
}
if !yield(*currentEvent, nil) {
return false
}
// Reset for next event.
currentEvent = nil
return true
}
// For fields that get their value replaced, we compare against the
// current total of all fields values rather than per-field, so
// re-assigned fields still count their prior length.
//
// This can reject some valid events but reliably bounds memory.
ensureRoom := func(addedSize int) bool {
if currentEvent == nil {
currentEvent = &Event{}
}
if currentEvent.len()+addedSize > MaxReadEventSize {
return false
}
return true
}
for scanner.Scan() {
line := scanner.Bytes()
if len(line) == 0 {
if !yieldEvent() {
return
}
continue
}
// Ignore "comment" lines.
//
// > If the line starts with a U+003A COLON character (:)
// > Ignore the line.
//
// https://html.spec.whatwg.org/multipage/server-sent-events.html#event-stream-interpretation
if line[0] == ':' {
continue
}
// A field line without colon is considered valid with empty value.
//
// > Note: If a line doesn't contain a colon, the entire line is treated as the field name with an empty value string.
//
// From https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events#fields
before, after, _ := bytes.Cut(line, []byte{':'})
// Per spec we must only trim the leading space from the field data.
// This avoids breaking the data sent by the server.
after = bytes.TrimPrefix(after, []byte{' '})
switch {
case bytes.Equal(before, sseEventKey):
if !ensureRoom(len(after)) {
yield(Event{}, trace.Wrap(ErrEventTooLarge))
return
}
currentEvent.Event = string(after)
case bytes.Equal(before, sseIDKey):
if !ensureRoom(len(after)) {
yield(Event{}, trace.Wrap(ErrEventTooLarge))
return
}
currentEvent.ID = string(after)
case bytes.Equal(before, sseRetryKey):
if !ensureRoom(len(after)) {
yield(Event{}, trace.Wrap(ErrEventTooLarge))
return
}
currentEvent.Retry = string(after)
case bytes.Equal(before, sseDataKey):
// Account for the new line added after the `data` field.
if !ensureRoom(len(after) + 1) {
yield(Event{}, trace.Wrap(ErrEventTooLarge))
return
}
// Cases where first data is empty, we cannot rely only on
// append as it would still return nil (so we wouldn't know a
// data field happened before). That's why we initialize the
// slice here.
if currentEvent.Data == nil {
currentEvent.Data = make([]byte, 0, len(after))
} else {
currentEvent.Data = append(currentEvent.Data, '\n')
}
currentEvent.Data = append(currentEvent.Data, after...)
default:
// Following the spec, unknown fields are just ignored.
continue
}
}
if err := scanner.Err(); err != nil {
if errors.Is(err, bufio.ErrTooLong) {
yield(Event{}, trace.Wrap(ErrEventTooLarge))
return
}
yield(Event{}, trace.Wrap(err))
return
}
yieldEvent()
}
}
// WriteEvent writes an (non-empty) SSE event into the provided Writer.
//
// This function does not flush the writer. Callers using HTTP streaming should
// flush their response after a successful write when events must be delivered
// promptly.
func WriteEvent(w io.Writer, event Event) (int, error) {
// Dropping empty events is acceptable as they don't carry values (only
// fields).
//
// In case callers need to write empty fields, we'd need a separate writer
// that preserves empty fields.
if event.Empty() {
return 0, nil
}
if err := event.validate(); err != nil {
return 0, trace.Wrap(err)
}
// Use internal buffer so we only write once to the target writer.
var buf bytes.Buffer
if event.Event != "" {
writeField(&buf, sseEventKey, event.Event)
}
if event.ID != "" {
writeField(&buf, sseIDKey, event.ID)
}
for data := range splitDataLines(event.Data) {
buf.Write(sseDataKey)
buf.WriteString(": ")
buf.Write(data)
buf.WriteByte('\n')
}
if event.Retry != "" {
writeField(&buf, sseRetryKey, event.Retry)
}
buf.WriteByte('\n')
n, err := buf.WriteTo(w)
return int(n), trace.Wrap(err)
}
// writeField internal helper that writes fields to buffer without allocating
// additional strings.
func writeField(b *bytes.Buffer, fieldName []byte, fieldData string) {
b.Write(fieldName)
b.WriteString(": ")
b.WriteString(fieldData)
b.WriteByte('\n')
}
// scanLines implements a custom scan line for bytes.Scanner that honors the
// SSE spec in addition to the [MaxReadEventSize].
//
// By the spec, the data can be split using \r (cr) \n (lf) following the end of
// line definition:
//
// end-of-line = ( cr lf / cr / lf )
//
// https://html.spec.whatwg.org/multipage/server-sent-events.html#parsing-an-event-stream
func scanLines() bufio.SplitFunc {
var skipLF bool
return func(data []byte, atEOF bool) (advance int, token []byte, err error) {
// If previous chunk ended on \r, swallow a leading \n here so \r\n
// count as one terminator across reads. Bare \r line endings still work.
if skipLF {
if len(data) == 0 {
return 0, nil, nil
}
skipLF = false
if data[0] == '\n' {
return 1, nil, nil
}
}
for i, b := range data {
if i > MaxReadEventSize {
return 0, nil, bufio.ErrTooLong
}
switch b {
case '\n':
return i + 1, data[:i], nil
case '\r':
// \r\n must count a single line break as per spec.
if i+1 < len(data) && data[i+1] == '\n' {
return i + 2, data[:i], nil
}
if i+1 == len(data) && !atEOF {
skipLF = true
}
return i + 1, data[:i], nil
}
}
if len(data) > MaxReadEventSize {
return 0, nil, bufio.ErrTooLong
}
if atEOF && len(data) > 0 {
return len(data), data, nil
}
return 0, nil, nil
}
}
// splitDataLines splits data into multiple event lines.
//
// By the spec, the data can be split using \r (cr) \n (lf) following the end of
// line definition:
//
// end-of-line = ( cr lf / cr / lf )
//
// https://html.spec.whatwg.org/multipage/server-sent-events.html#parsing-an-event-stream
func splitDataLines(data []byte) iter.Seq[[]byte] {
return func(yield func([]byte) bool) {
// No data available, return nothing to the caller.
if len(data) == 0 {
return
}
for {
i := bytes.IndexAny(data, "\r\n")
if i < 0 {
yield(data)
return
}
if !yield(data[:i]) {
return
}
// \r\n must count a single line break as per spec.
if data[i] == '\r' && i+1 < len(data) && data[i+1] == '\n' {
data = data[i+2:]
} else {
data = data[i+1:]
}
}
}
}
func onlyDigits(str string) bool {
return strings.IndexFunc(str, func(r rune) bool {
return r < '0' || r > '9'
}) == -1
}
// SSE supported fields names.
//
// Ref: https://developer.mozilla.org/en-US/docs/Web/API/Server-sent_events/Using_server-sent_events#fields.
var (
sseEventKey = []byte("event")
sseIDKey = []byte("id")
sseDataKey = []byte("data")
sseRetryKey = []byte("retry")
)
/*
* 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/>.
*/
//nolint:goimports // goimports disagree with gci on blank imports
package proxy
import (
"context"
"crypto/tls"
"fmt"
"log/slog"
"net"
"net/http"
"net/url"
"github.com/gravitational/trace"
authzapi "k8s.io/api/authorization/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
utilnet "k8s.io/apimachinery/pkg/util/net"
"k8s.io/client-go/kubernetes"
authztypes "k8s.io/client-go/kubernetes/typed/authorization/v1"
// Load kubeconfig auth plugins for gcp and azure.
// Without this, users can't provide a kubeconfig using those.
//
// Note: we don't want to load _all_ plugins. This is a balance between
// support for popular hosting providers and minimizing attack surface.
_ "k8s.io/client-go/plugin/pkg/client/auth/azure"
_ "k8s.io/client-go/plugin/pkg/client/auth/gcp"
"k8s.io/client-go/rest"
"k8s.io/client-go/transport"
"github.com/gravitational/teleport/api/types"
kubeutils "github.com/gravitational/teleport/lib/kube/utils"
"github.com/gravitational/teleport/lib/service/servicecfg"
)
// getKubeDetails fetches the kubernetes API credentials.
//
// There are 2 possible sources of credentials:
// - pod service account credentials: files in hardcoded paths when running
// inside of a k8s pod; this is used when kubeClusterName is set
// - kubeconfig: a file with a set of k8s endpoints and credentials mapped to
// them this is used when kubeconfigPath is set
//
// serviceType changes the loading behavior:
// - LegacyProxyService:
// - if loading from kubeconfig, only "current-context" is returned; the
// returned map key matches tpClusterName
// - if no credentials are loaded, no error is returned
// - permission self-test failures are only logged
//
// - ProxyService:
// - no credentials are loaded and no error is returned
//
// - KubeService:
// - if loading from kubeconfig, all contexts are returned
// - if no credentials are loaded, returns an error
// - permission self-test failures cause an error to be returned
func (f *Forwarder) getKubeDetails(ctx context.Context) error {
serviceType := f.cfg.KubeServiceType
kubeconfigPath := f.cfg.KubeconfigPath
kubeClusterName := f.cfg.KubeClusterName
tpClusterName := f.cfg.ClusterName
f.log.DebugContext(ctx, "Reading Kubernetes details",
"kubeconfig_path", kubeconfigPath,
"kube_cluster_name", kubeClusterName,
"service_type", serviceType,
)
// Proxy service should never have creds, forwards to kube service
if serviceType == ProxyService {
return nil
}
// Load kubeconfig or local pod credentials.
loadAll := serviceType == KubeService
cfg, err := kubeutils.GetKubeConfig(kubeconfigPath, loadAll, kubeClusterName)
if err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err)
}
if trace.IsNotFound(err) || len(cfg.Contexts) == 0 {
switch serviceType {
case KubeService:
return trace.BadParameter("no Kubernetes credentials found; Kubernetes_service requires either a valid kubeconfig_file or to run inside of a Kubernetes pod")
case LegacyProxyService:
f.log.DebugContext(ctx, "Could not load Kubernetes credentials. This proxy will still handle Kubernetes requests for trusted teleport clusters or Kubernetes nodes in this teleport cluster")
}
return nil
}
if serviceType == LegacyProxyService {
// Hack for legacy proxy service - register a k8s cluster named after
// the teleport cluster name to route legacy requests.
//
// Also, remove all other contexts. Multiple kubeconfig entries are
// only supported for kubernetes_service.
if currentContext, ok := cfg.Contexts[cfg.CurrentContext]; ok {
cfg.Contexts = map[string]*rest.Config{
tpClusterName: currentContext,
}
} else {
return trace.BadParameter("no Kubernetes current-context found; Kubernetes proxy service requires either a valid kubeconfig_file with a current-context or to run inside of a Kubernetes pod")
}
}
// Convert kubeconfig contexts into kubeCreds.
for cluster, clientCfg := range cfg.Contexts {
clusterCreds, err := extractKubeCreds(ctx, serviceType, cluster, clientCfg, f.log, f.cfg.CheckImpersonationPermissions)
if err != nil {
f.log.WarnContext(ctx, "failed to load credentials for cluster",
"cluster", cluster,
"error", err,
)
continue
}
kubeCluster, err := types.NewKubernetesClusterV3(
types.Metadata{
Name: cluster,
}, types.KubernetesClusterSpecV3{},
types.KubeClusterWithScope(f.cfg.GetScope()),
)
if err != nil {
f.log.WarnContext(ctx, "failed to create KubernetesClusterV3 from credentials for cluster",
"cluster", cluster,
"error", err,
)
continue
}
details, err := newClusterDetails(ctx,
clusterDetailsConfig{
cluster: kubeCluster,
kubeCreds: clusterCreds,
log: f.log.With("cluster", kubeCluster.GetName()),
checker: f.cfg.CheckImpersonationPermissions,
component: serviceType,
clock: f.cfg.Clock,
})
if err != nil {
f.log.WarnContext(ctx, "Failed to create cluster details for cluster",
"cluster", cluster,
"error", err,
)
return trace.Wrap(err)
}
f.clusterDetails[cluster] = details
}
return nil
}
func extractKubeCreds(ctx context.Context, component string, cluster string, clientCfg *rest.Config, log *slog.Logger, checkPermissions servicecfg.ImpersonationPermissionsChecker) (*staticKubeCreds, error) {
log = log.With("cluster", cluster)
// Disable client-go's client-side rate limiter (default 5 QPS).
// The kube exec path fetches pod metadata through this client,
// so the default limit caps exec throughput at ~5 QPS per agent.
clientCfg.QPS = -1
log.DebugContext(ctx, "Checking Kubernetes impersonation permissions")
client, err := kubernetes.NewForConfig(clientCfg)
if err != nil {
return nil, trace.Wrap(err, "failed to generate Kubernetes client for cluster %q", cluster)
}
// For each loaded cluster, check impersonation permissions. This
// check only logs when permissions are not configured, but does not fail startup.
if err := checkPermissions(ctx, cluster, client.AuthorizationV1().SelfSubjectAccessReviews()); err != nil {
log.WarnContext(ctx, "Failed to test the necessary Kubernetes permissions. The target Kubernetes cluster may be down or have misconfigured RBAC. This teleport instance will still handle Kubernetes requests towards this Kubernetes cluster.",
"error", err,
)
} else {
log.DebugContext(ctx, "Have all necessary Kubernetes impersonation permissions")
}
targetAddr, err := parseKubeHost(clientCfg.Host)
if err != nil {
return nil, trace.Wrap(err)
}
// tlsConfig can be nil and still no error is returned.
// This happens when no `certificate-authority-data` is provided in kubeconfig because one is expected to use
// the system default CA pool.
tlsConfig, err := rest.TLSConfigFor(clientCfg)
if err != nil {
return nil, trace.Wrap(err, "failed to generate TLS config from kubeconfig: %v", err)
}
transportConfig, err := clientCfg.TransportConfig()
if err != nil {
return nil, trace.Wrap(err, "failed to generate transport config from kubeconfig: %v", err)
}
transport, err := newDirectTransport(component, tlsConfig, transportConfig)
if err != nil {
return nil, trace.Wrap(err, "failed to generate transport from kubeconfig: %v", err)
}
log.DebugContext(ctx, "Initialized Kubernetes credentials")
return &staticKubeCreds{
tlsConfig: tlsConfig,
transportConfig: transportConfig,
targetAddr: targetAddr,
kubeClient: client,
clientRestCfg: clientCfg,
transport: transport,
}, nil
}
// newDirectTransport creates a new http.Transport that will be used to connect to the Kubernetes API server.
// It is a direct connection, not going through a Teleport proxy.
// The transport used respects HTTP_PROXY, HTTPS_PROXY, and NO_PROXY environment variables.
func newDirectTransport(component string, tlsConfig *tls.Config, transportConfig *transport.Config) (http.RoundTripper, error) {
h2HTTPTransport, err := newH2Transport(tlsConfig, nil)
if err != nil {
return nil, trace.Wrap(err)
}
// SetTransportDefaults sets the default values for the transport including
// support for HTTP_PROXY, HTTPS_PROXY, NO_PROXY, and the default user agent.
h2HTTPTransport = utilnet.SetTransportDefaults(h2HTTPTransport)
h2Transport, err := wrapTransport(h2HTTPTransport, transportConfig)
if err != nil {
return nil, trace.Wrap(err)
}
return instrumentedRoundtripper(component, h2Transport), nil
}
// parseKubeHost parses and formats kubernetes hostname
// to host:port format, if no port it set,
// it assumes default HTTPS port
func parseKubeHost(host string) (string, error) {
u, err := url.Parse(host)
if err != nil {
return "", trace.Wrap(err, "failed to parse Kubernetes host: %v", err)
}
if _, _, err := net.SplitHostPort(u.Host); err != nil {
// add default HTTPS port
return fmt.Sprintf("%v:443", u.Host), nil
}
return u.Host, nil
}
func checkImpersonationPermissions(ctx context.Context, cluster string, sarClient authztypes.SelfSubjectAccessReviewInterface) error {
for _, resource := range []string{"users", "groups", "serviceaccounts"} {
resp, err := sarClient.Create(ctx, &authzapi.SelfSubjectAccessReview{
Spec: authzapi.SelfSubjectAccessReviewSpec{
ResourceAttributes: &authzapi.ResourceAttributes{
Verb: "impersonate",
Resource: resource,
},
},
}, metav1.CreateOptions{})
if err != nil {
return trace.Wrap(err, "failed to verify impersonation permissions for Kubernetes: %v; this may be due to missing the SelfSubjectAccessReview permission on the ClusterRole used by the proxy; please make sure that proxy has all the necessary permissions: https://goteleport.com/docs/enroll-resources/kubernetes-access/controls/#enabling-impersonation", err)
}
if !resp.Status.Allowed {
return trace.AccessDenied("proxy can't impersonate Kubernetes %s at the cluster level; please make sure that proxy has all the necessary permissions: https://goteleport.com/docs/enroll-resources/kubernetes-access/controls/#enabling-impersonation", resource)
}
}
return nil
}
/*
* 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 proxy
import (
"context"
"encoding/base64"
"log/slog"
"maps"
"net/http"
"strings"
"sync"
"time"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/service/eks"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"golang.org/x/sync/singleflight"
authzapi "k8s.io/api/authorization/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime/schema"
"k8s.io/apimachinery/pkg/runtime/serializer"
"k8s.io/apimachinery/pkg/version"
"k8s.io/client-go/rest"
"k8s.io/client-go/tools/clientcmd"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/retryutils"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/cloud/azure"
"github.com/gravitational/teleport/lib/cloud/gcp"
kubeutils "github.com/gravitational/teleport/lib/kube/utils"
"github.com/gravitational/teleport/lib/labels"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/services"
)
// kubeDetails contain the cluster-related details including authentication.
type kubeDetails struct {
kubeCreds
// dynamicLabels is the dynamic labels executor for this cluster.
dynamicLabels *labels.Dynamic
// kubeCluster is the dynamic kube_cluster or a static generated from kubeconfig and that only has the name populated.
kubeCluster types.KubeCluster
// kubeClusterVersion is the version of the kube_cluster's related Kubernetes server.
kubeClusterVersion *version.Info
// rwMu is the mutex to protect the kubeCodecs, gvkSupportedResources, and rbacSupportedTypes.
rwMu sync.RWMutex
// kubeCodecs is the codec factory for the cluster resources.
// The codec factory includes the default resources and the namespaced resources
// that are supported by the cluster.
// The codec factory is updated periodically to include the latest custom resources
// that are added to the cluster.
kubeCodecs *serializer.CodecFactory
// rbacSupportedTypes is the list of supported types for RBAC for the cluster.
// The list is updated periodically to include the latest custom resources
// that are added to the cluster.
rbacSupportedTypes rbacSupportedResources
// gvkSupportedResources is the list of registered API path resources and their
// GVK definition.
gvkSupportedResources gvkSupportedResources
// refreshGroup is used to coalesce concurrent discovery of the same API group-version into one request.
refreshGroup singleflight.Group
// isClusterOffline is true if the cluster is offline.
// An offline cluster will not be able to serve any requests until it comes back online.
// The cluster is marked as offline if the cluster schema cannot be created
// and the list of supported types for RBAC cannot be generated.
isClusterOffline bool
cancelFunc context.CancelFunc
wg sync.WaitGroup
}
// clusterDetailsConfig contains the configuration for creating a proxied cluster.
type clusterDetailsConfig struct {
// azureClients provides Azure SDK clients
azureClients azure.Clients
// gcpClients provides GCP SDK clients
gcpClients gcp.Clients
// awsCloudClients provides AWS SDK clients.
awsCloudClients AWSClientGetter
// kubeCreds is the credentials to use for the cluster.
kubeCreds kubeCreds
// cluster is the cluster to create a proxied cluster for.
cluster types.KubeCluster
// log is the logger to use.
log *slog.Logger
// checker is the permissions checker to use.
checker servicecfg.ImpersonationPermissionsChecker
// resourceMatchers is the list of resource matchers to match the cluster against
// to determine if we should assume the role or not for AWS.
resourceMatchers []services.ResourceMatcher
// clock is the clock to use.
clock clockwork.Clock
// component is the Kubernetes component that serves this cluster.
component KubeServiceType
}
const (
defaultRefreshPeriod = 5 * time.Minute
backoffRefreshStep = 10 * time.Second
)
// newClusterDetails creates a proxied kubeDetails structure given a dynamic cluster.
func newClusterDetails(ctx context.Context, cfg clusterDetailsConfig) (_ *kubeDetails, err error) {
creds := cfg.kubeCreds
if creds == nil {
creds, err = getKubeClusterCredentials(ctx, cfg)
if err != nil {
return nil, trace.Wrap(err)
}
}
var dynLabels *labels.Dynamic
if len(cfg.cluster.GetDynamicLabels()) > 0 {
dynLabels, err = labels.NewDynamic(
ctx,
&labels.DynamicConfig{
Labels: cfg.cluster.GetDynamicLabels(),
Log: cfg.log,
})
if err != nil {
return nil, trace.Wrap(err)
}
dynLabels.Sync()
go dynLabels.Start()
}
var isClusterOffline bool
// Create the codec factory and the list of supported types for RBAC.
codecFactory, rbacSupportedTypes, gvkSupportedRes, err := newClusterSchemaBuilder(cfg.log, creds.getKubeClient())
if err != nil {
cfg.log.WarnContext(ctx, "Failed to create cluster schema, the cluster may be offline", "error", err)
// If the cluster is offline, we will not be able to create the codec factory
// and the list of supported types for RBAC.
// We mark the cluster as offline and continue to create the kubeDetails but
// the offline cluster will not be able to serve any requests until it comes back online.
isClusterOffline = true
}
kubeVersion, err := creds.getKubeClient().Discovery().ServerVersion()
if err != nil {
cfg.log.WarnContext(ctx, "Failed to get Kubernetes cluster version, the cluster may be offline", "error", err)
}
ctx, cancel := context.WithCancel(ctx)
k := &kubeDetails{
kubeCreds: creds,
dynamicLabels: dynLabels,
kubeCluster: cfg.cluster,
kubeClusterVersion: kubeVersion,
kubeCodecs: codecFactory,
rbacSupportedTypes: rbacSupportedTypes,
cancelFunc: cancel,
isClusterOffline: isClusterOffline,
gvkSupportedResources: gvkSupportedRes,
}
// If cluster is online and there's no errors, we refresh details seldom (every 5 minutes),
// but if the cluster is offline, we try to refresh details more often to catch it getting back online earlier.
firstPeriod := defaultRefreshPeriod
if isClusterOffline {
firstPeriod = backoffRefreshStep
}
refreshDelay, err := retryutils.NewLinear(retryutils.LinearConfig{
First: firstPeriod,
Step: backoffRefreshStep,
Max: defaultRefreshPeriod,
Jitter: retryutils.SeventhJitter,
Clock: cfg.clock,
})
if err != nil {
k.Close()
return nil, trace.Wrap(err)
}
k.wg.Add(1)
// Start the periodic update of the codec factory and the list of supported types for RBAC.
go func() {
defer k.wg.Done()
for {
select {
case <-ctx.Done():
return
case <-refreshDelay.After():
codecFactory, rbacSupportedTypes, gvkSupportedResources, err := newClusterSchemaBuilder(cfg.log, creds.getKubeClient())
if err != nil {
// If this is first time we get an error, we reset retry mechanism so it will start trying to refresh details quicker, with linear backoff.
if refreshDelay.First == defaultRefreshPeriod {
refreshDelay.First = backoffRefreshStep
refreshDelay.Reset()
} else {
refreshDelay.Inc()
}
cfg.log.ErrorContext(ctx, "Failed to update cluster schema", "error", err)
continue
}
kubeVersion, err := creds.getKubeClient().Discovery().ServerVersion()
if err != nil {
cfg.log.WarnContext(ctx, "Failed to get Kubernetes cluster version, the cluster may be offline", "error", err)
}
// Restore details refresh delay to the default value, in case previously cluster was offline.
refreshDelay.First = defaultRefreshPeriod
k.rwMu.Lock()
k.kubeCodecs = codecFactory
k.rbacSupportedTypes = rbacSupportedTypes
k.gvkSupportedResources = gvkSupportedResources
k.isClusterOffline = false
k.kubeClusterVersion = kubeVersion
k.rwMu.Unlock()
}
}
}()
return k, nil
}
func (k *kubeDetails) Close() {
// send a close signal and wait for the close to finish.
k.cancelFunc()
k.wg.Wait()
if k.dynamicLabels != nil {
k.dynamicLabels.Close()
}
// it is safe to call close even for static creds.
k.kubeCreds.close()
}
// getClusterSupportedResources returns the codec factory and the list of supported types for RBAC.
func (k *kubeDetails) getClusterSupportedResources() (*serializer.CodecFactory, rbacSupportedResources, error) {
k.rwMu.RLock()
defer k.rwMu.RUnlock()
// If the cluster is offline, return an error because we don't have the schema
// for the cluster.
if k.isClusterOffline {
return nil, nil, trace.ConnectionProblem(nil, "kubernetes cluster %q is offline", k.kubeCluster.GetName())
}
return k.kubeCodecs, k.rbacSupportedTypes, nil
}
// resolveResource looks up a resource kind in the discovery-backed cache.
// On a miss it discovers just that API group-version (a single request) to catch recently-installed CRDs,
// then retries, and returns the definition and whether it was found.
func (k *kubeDetails) resolveResource(apiGroup, apiGroupVersion, resourceKind string) (res metav1.APIResource, found bool) {
k.rwMu.RLock()
res, found = k.rbacSupportedTypes.getResource(apiGroup, resourceKind)
k.rwMu.RUnlock()
if found {
return res, true
}
k.refreshGroupVersion(apiGroup, apiGroupVersion)
k.rwMu.RLock()
defer k.rwMu.RUnlock()
return k.rbacSupportedTypes.getResource(apiGroup, resourceKind)
}
// refreshGroupVersion discovers a single API group-version and merges any new resources into the cache.
// When the version is empty (SelfSubjectAccessReview requests often omit it) it resolves the group's
// preferred version first. Concurrent calls for the same group(-version) are coalesced — including the
// preferred-version lookup — so version-less requests don't each hit discovery. A discovery error leaves
// the kind absent (and thus denied).
func (k *kubeDetails) refreshGroupVersion(apiGroup, apiGroupVersion string) {
if k.kubeCreds == nil {
return
}
_, _, _ = k.refreshGroup.Do(apiGroup+"/"+apiGroupVersion, func() (any, error) {
version := apiGroupVersion
if version == "" {
if version = k.preferredVersionForGroup(apiGroup); version == "" {
return nil, nil
}
}
gv := schema.GroupVersion{Group: apiGroup, Version: version}
list, err := k.getKubeClient().Discovery().ServerResourcesForGroupVersion(gv.String())
if err != nil {
return nil, nil
}
k.mergeGroupVersion(gv, list)
return nil, nil
})
}
// preferredVersionForGroup returns the cluster's preferred version for an API group, or "" if the
// group can't be resolved. Lets resolveResource discover a group whose version the caller didn't supply.
func (k *kubeDetails) preferredVersionForGroup(apiGroup string) string {
groups, err := k.getKubeClient().Discovery().ServerGroups()
if err != nil {
return ""
}
for _, g := range groups.Groups {
if g.Name == apiGroup {
return g.PreferredVersion.Version
}
}
return ""
}
// mergeGroupVersion merges a group-version's resources into the cached RBAC types, GVK map, and codec scheme,
// swapping in fresh copies so existing readers keep a consistent snapshot.
// It only rebuilds when there are new kinds.
func (k *kubeDetails) mergeGroupVersion(gv schema.GroupVersion, list *metav1.APIResourceList) {
k.rwMu.Lock()
defer k.rwMu.Unlock()
hasNew := false
for _, apiResource := range list.APIResources {
if _, ok := k.rbacSupportedTypes[allowedResourcesKey{apiGroup: gv.Group, resourceKind: apiResource.Name}]; !ok {
hasNew = true
break
}
}
if !hasNew {
return
}
rbac := maps.Clone(k.rbacSupportedTypes)
gvk := maps.Clone(k.gvkSupportedResources)
for _, apiResource := range list.APIResources {
apiResource.Group = gv.Group
apiResource.Version = gv.Version
rbac[allowedResourcesKey{apiGroup: gv.Group, resourceKind: apiResource.Name}] = apiResource
gvk[gvkSupportedResourcesKey{name: apiResource.Name, apiGroup: gv.Group, version: gv.Version}] = &schema.GroupVersionKind{
Group: gv.Group,
Version: gv.Version,
Kind: apiResource.Kind,
}
}
codecs, err := buildCodecsForGVKs(gvk)
if err != nil {
return
}
k.rbacSupportedTypes = rbac
k.gvkSupportedResources = gvk
k.kubeCodecs = codecs
}
// getObjectGVK returns the default GVK (if any) registered for the specified request path.
func (k *kubeDetails) getObjectGVK(resource apiResource) *schema.GroupVersionKind {
k.rwMu.RLock()
defer k.rwMu.RUnlock()
return k.gvkSupportedResources[gvkSupportedResourcesKey{
name: strings.Split(resource.resourceKind, "/")[0],
apiGroup: resource.apiGroup,
version: resource.apiGroupVersion,
}]
}
// GetProtocol returns the network protocol used for checking health.
func (t *kubeDetails) GetProtocol() types.TargetHealthProtocol {
return types.TargetHealthProtocolHTTP
}
type operation struct {
verb string
resource string
display string
}
var permissionOps = []operation{
{
verb: "impersonate",
resource: "users",
display: "impersonate users",
},
{
verb: "impersonate",
resource: "groups",
display: "impersonate groups",
},
{
verb: "impersonate",
resource: "serviceaccounts",
display: "impersonate service accounts",
},
{
verb: "get",
resource: "pods",
display: "get pods",
},
}
const errorGuide = "Please see the Kubernetes Access Troubleshooting guide, https://goteleport.com/docs/enroll-resources/kubernetes-access/troubleshooting."
// CheckHealth checks the health of a Kubernetes cluster.
func (k *kubeDetails) CheckHealth(ctx context.Context) ([]string, error) {
addresses := []string{k.getTargetAddr()}
client := k.getKubeClient().AuthorizationV1().SelfSubjectAccessReviews()
// Check permissions to the Kubernetes cluster.
var missingPermissions []string
for _, op := range permissionOps {
resp, err := client.Create(ctx, &authzapi.SelfSubjectAccessReview{
Spec: authzapi.SelfSubjectAccessReviewSpec{
ResourceAttributes: &authzapi.ResourceAttributes{
Verb: op.verb,
Resource: op.resource,
},
},
}, metav1.CreateOptions{})
if err != nil {
// Check whether the Kubernetes cluster is down.
// Avoid reporting permissions errors when the Kubernetes cluster is down.
if readyzErr := k.checkHealthReadyz(ctx); readyzErr != nil {
return addresses, trace.Wrap(readyzErr)
}
return addresses, trace.Wrap(err, "Unable to check Kubernetes permissions. %s", errorGuide)
}
if !resp.Status.Allowed {
missingPermissions = append(missingPermissions, op.display)
}
}
if len(missingPermissions) > 0 {
return addresses, trace.AccessDenied("Missing required Kubernetes permissions: %s. %s",
strings.Join(missingPermissions, ", "),
errorGuide)
}
return addresses, nil
}
// checkHealthReadyz checks the health of a Kubernetes cluster with the `/readyz` endpoint.
func (k *kubeDetails) checkHealthReadyz(ctx context.Context) error {
readyzResult := k.getKubeClient().Discovery().RESTClient().Get().AbsPath("/readyz").Do(ctx)
if err := readyzResult.Error(); err != nil {
return trace.ConnectionProblem(err, "Unable to contact the Kubernetes cluster. %s", errorGuide)
}
var statusCode int
readyzResult.StatusCode(&statusCode)
if statusCode != http.StatusOK {
return trace.ConnectionProblem(nil, "Unhealthy Kubernetes cluster detected with status code %d. %s", statusCode, errorGuide)
}
return nil
}
// getKubeClusterCredentials generates kube credentials for dynamic clusters.
func getKubeClusterCredentials(ctx context.Context, cfg clusterDetailsConfig) (kubeCreds, error) {
switch dynCredsCfg := (dynamicCredsConfig{
kubeCluster: cfg.cluster,
log: cfg.log,
checker: cfg.checker,
resourceMatchers: cfg.resourceMatchers,
clock: cfg.clock,
component: cfg.component,
}); {
case cfg.cluster.IsKubeconfig():
return getStaticCredentialsFromKubeconfig(ctx, cfg.component, cfg.cluster, cfg.log, cfg.checker)
case cfg.cluster.IsAzure():
return getAzureCredentials(ctx, cfg.azureClients, dynCredsCfg)
case cfg.cluster.IsAWS():
return getAWSCredentials(ctx, cfg.awsCloudClients, dynCredsCfg)
case cfg.cluster.IsGCP():
return getGCPCredentials(ctx, cfg.gcpClients, dynCredsCfg)
default:
return nil, trace.BadParameter("authentication method provided for cluster %q not supported", cfg.cluster.GetName())
}
}
// getAzureCredentials creates a dynamicCreds that generates and updates the access credentials to a AKS Kubernetes cluster.
func getAzureCredentials(ctx context.Context, azureClients azure.Clients, cfg dynamicCredsConfig) (*dynamicKubeCreds, error) {
// create a client that returns the credentials for kubeCluster
cfg.client = azureRestConfigClient(azureClients)
creds, err := newDynamicKubeCreds(
ctx,
cfg,
)
return creds, trace.Wrap(err)
}
// azureRestConfigClient creates a dynamicCredsClient that returns credentials to a AKS cluster.
func azureRestConfigClient(azureClients azure.Clients) dynamicCredsClient {
return func(ctx context.Context, cluster types.KubeCluster) (*rest.Config, time.Time, error) {
aksClient, err := azureClients.GetKubernetesClient(ctx, cluster.GetAzureConfig().SubscriptionID)
if err != nil {
return nil, time.Time{}, trace.Wrap(err)
}
cfg, exp, err := aksClient.ClusterCredentials(ctx, azure.ClusterCredentialsConfig{
ResourceGroup: cluster.GetAzureConfig().ResourceGroup,
ResourceName: cluster.GetAzureConfig().ResourceName,
TenantID: cluster.GetAzureConfig().TenantID,
ImpersonationPermissionsChecker: checkImpersonationPermissions,
})
return cfg, exp, trace.Wrap(err)
}
}
// getAWSCredentials creates a dynamicKubeCreds that generates and updates the access credentials to a EKS kubernetes cluster.
func getAWSCredentials(ctx context.Context, cloudClients AWSClientGetter, cfg dynamicCredsConfig) (*dynamicKubeCreds, error) {
// create a client that returns the credentials for kubeCluster
cfg.client = getAWSClientRestConfig(cloudClients, cfg.clock, cfg.resourceMatchers)
creds, err := newDynamicKubeCreds(ctx, cfg)
return creds, trace.Wrap(err)
}
// getAWSResourceMatcherToCluster returns the AWS assume role ARN and external ID for the cluster that matches the kubeCluster.
// If no match is found, nil is returned, which means that we should not attempt to assume a role.
func getAWSResourceMatcherToCluster(kubeCluster types.KubeCluster, resourceMatchers []services.ResourceMatcher) *services.ResourceMatcherAWS {
if !kubeCluster.IsAWS() {
return nil
}
for _, matcher := range resourceMatchers {
if len(matcher.Labels) == 0 || matcher.AWS.AssumeRoleARN == "" {
continue
}
if match, _, _ := services.MatchLabels(matcher.Labels, kubeCluster.GetAllLabels()); !match {
continue
}
return &matcher.AWS
}
return nil
}
// STSPresignClient is the subset of the STS presign interface we use in fetchers.
type STSPresignClient = kubeutils.STSPresignClient
// EKSClient is the subset of the EKS Client interface we use.
type EKSClient interface {
eks.DescribeClusterAPIClient
}
// AWSClientGetter is an interface for getting an EKS client and an STS client.
type AWSClientGetter interface {
awsconfig.Provider
// GetAWSEKSClient returns AWS EKS client for the specified config.
GetAWSEKSClient(aws.Config) EKSClient
// GetAWSSTSPresignClient returns AWS STS presign client for the specified config.
GetAWSSTSPresignClient(aws.Config) STSPresignClient
}
// getAWSClientRestConfig creates a dynamicCredsClient that generates returns credentials to EKS clusters.
func getAWSClientRestConfig(cloudClients AWSClientGetter, clock clockwork.Clock, resourceMatchers []services.ResourceMatcher) dynamicCredsClient {
return func(ctx context.Context, cluster types.KubeCluster) (*rest.Config, time.Time, error) {
region := cluster.GetAWSConfig().Region
opts := []awsconfig.OptionsFn{
awsconfig.WithAmbientCredentials(),
}
if awsAssume := getAWSResourceMatcherToCluster(cluster, resourceMatchers); awsAssume != nil {
opts = append(opts, awsconfig.WithAssumeRole(awsAssume.AssumeRoleARN, awsAssume.ExternalID))
}
cfg, err := cloudClients.GetConfig(ctx, region, opts...)
if err != nil {
return nil, time.Time{}, trace.Wrap(err)
}
regionalClient := cloudClients.GetAWSEKSClient(cfg)
eksCfg, err := regionalClient.DescribeCluster(ctx, &eks.DescribeClusterInput{
Name: aws.String(cluster.GetAWSConfig().Name),
})
if err != nil {
return nil, time.Time{}, trace.Wrap(err)
}
ca, err := base64.StdEncoding.DecodeString(aws.ToString(eksCfg.Cluster.CertificateAuthority.Data))
if err != nil {
return nil, time.Time{}, trace.Wrap(err)
}
apiEndpoint := aws.ToString(eksCfg.Cluster.Endpoint)
if len(apiEndpoint) == 0 {
return nil, time.Time{}, trace.BadParameter("invalid api endpoint for cluster %q", cluster.GetAWSConfig().Name)
}
stsPresignClient := cloudClients.GetAWSSTSPresignClient(cfg)
token, exp, err := kubeutils.GenAWSEKSToken(ctx, stsPresignClient, cluster.GetAWSConfig().Name, clock)
if err != nil {
return nil, time.Time{}, trace.Wrap(err)
}
return &rest.Config{
Host: apiEndpoint,
BearerToken: token,
TLSClientConfig: rest.TLSClientConfig{
CAData: ca,
},
}, exp, nil
}
}
// getStaticCredentialsFromKubeconfig loads a kubeconfig from the cluster and returns the access credentials for the cluster.
// If the config defines multiple contexts, it will pick one (the order is not guaranteed).
func getStaticCredentialsFromKubeconfig(ctx context.Context, component KubeServiceType, cluster types.KubeCluster, log *slog.Logger, checker servicecfg.ImpersonationPermissionsChecker) (*staticKubeCreds, error) {
config, err := clientcmd.Load(cluster.GetKubeconfig())
if err != nil {
return nil, trace.WrapWithMessage(err, "unable to parse kubeconfig for cluster %q", cluster.GetName())
}
if len(config.CurrentContext) == 0 && len(config.Contexts) > 0 {
// select the first context key as default context
for k := range config.Contexts {
config.CurrentContext = k
break
}
}
restConfig, err := clientcmd.NewDefaultClientConfig(*config, nil).ClientConfig()
if err != nil {
return nil, trace.WrapWithMessage(err, "unable to create client from kubeconfig for cluster %q", cluster.GetName())
}
creds, err := extractKubeCreds(ctx, component, cluster.GetName(), restConfig, log, checker)
return creds, trace.Wrap(err)
}
// getGCPCredentials creates a dynamicKubeCreds that generates and updates the access credentials to a GKE kubernetes cluster.
func getGCPCredentials(ctx context.Context, gcpClients gcp.Clients, cfg dynamicCredsConfig) (*dynamicKubeCreds, error) {
// create a client that returns the credentials for kubeCluster
cfg.client = gcpRestConfigClient(gcpClients)
creds, err := newDynamicKubeCreds(ctx, cfg)
return creds, trace.Wrap(err)
}
// gcpRestConfigClient creates a dynamicCredsClient that returns credentials to a GKE cluster.
func gcpRestConfigClient(gcpClients gcp.Clients) dynamicCredsClient {
return func(ctx context.Context, cluster types.KubeCluster) (*rest.Config, time.Time, error) {
gkeClient, err := gcpClients.GetGKEClient(ctx)
if err != nil {
return nil, time.Time{}, trace.Wrap(err)
}
cfg, exp, err := gkeClient.GetClusterRestConfig(ctx,
gcp.ClusterDetails{
ProjectID: cluster.GetGCPConfig().ProjectID,
Location: cluster.GetGCPConfig().Location,
Name: cluster.GetGCPConfig().Name,
},
)
return cfg, exp, trace.Wrap(err)
}
}
/*
* 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 proxy
import (
"compress/gzip"
"io"
"net/http"
"sync"
"github.com/gravitational/trace"
)
const (
// contentEncodingHeader is the HTTP header used to specify the
// content encoding of the response.
contentEncodingHeader = "Content-Encoding"
// contentEncodingGZIP is the value for the Content-Encoding header when
// the response is compressed with gzip.
contentEncodingGZIP = "gzip"
// defaultGzipContentEncodingLevel is set to 1 which uses least CPU compared to higher levels, yet offers
// similar compression ratios (off by at most 1.5x, but typically within 1.1x-1.3x). For further details see -
// https://github.com/kubernetes/kubernetes/issues/112296
defaultGzipContentEncodingLevel = 1
)
var gzipPool = &sync.Pool{
New: func() any {
gw, err := gzip.NewWriterLevel(nil, defaultGzipContentEncodingLevel)
if err != nil {
// This should never happen.
panic(err)
}
return gw
},
}
type (
// compressionFunc is a function that decompresses data.
decompressionFunc func(dst io.Writer, src io.Reader) error
// compressionFunc is a function that returns a WriteCloser that compresses data
// written to it into the provided io.Writer.
compressionFunc func(dst io.Writer) io.WriteCloser
)
// getResponseCompressorDecompressor returns a compression and decompression function based on the
// Content-Encoding header.
func getResponseCompressorDecompressor(headers http.Header) (compressor compressionFunc, decompressor decompressionFunc, err error) {
encoding := headers.Get(contentEncodingHeader)
switch encoding {
case contentEncodingGZIP:
compressor = func(dst io.Writer) io.WriteCloser {
gzw := gzipPool.Get().(*gzip.Writer)
gzw.Reset(dst)
return &gzipWrapper{gzw}
}
decompressor = func(dst io.Writer, src io.Reader) error {
gzr, err := gzip.NewReader(src)
if err != nil {
return trace.Wrap(err)
}
defer gzr.Close()
_, err = io.Copy(dst, gzr)
return trace.Wrap(err)
}
return
case "":
compressor = func(dst io.Writer) io.WriteCloser {
return &nopCloserWrapper{dst}
}
decompressor = func(dst io.Writer, src io.Reader) error {
_, err := io.Copy(dst, src)
return trace.Wrap(err)
}
return
default:
return nil, nil, trace.BadParameter("unknown encoding %q", encoding)
}
}
// gzipWrapper wraps a gzip.Writer to implement io.WriteCloser.
// When Close is called, the underlying gzip.Writer is returned to the pool.
type gzipWrapper struct {
*gzip.Writer
}
// Close closes the underlying writter and returns it to the pool.
func (w *gzipWrapper) Close() error {
err := w.Writer.Close()
w.Writer.Reset(nil)
gzipPool.Put(w.Writer)
w.Writer = nil
return trace.Wrap(err)
}
// nopCloserWrapper wraps an io.Writer to implement io.WriteCloser.
type nopCloserWrapper struct {
io.Writer
}
// Close has no action on the underlying writer.
func (*nopCloserWrapper) Close() error {
return nil
}
// wrapContentEncoding returns a reader/writer pair that handles Content-Encoding.
// For gzip, it wraps the reader in a gzip decompressor and the writer in a pooled gzip compressor.
// For identity or no encoding, it returns no-op closers.
// Returns an error for unsupported encodings.
func wrapContentEncoding(r io.Reader, w io.Writer, contentEncoding string) (io.ReadCloser, io.WriteCloser, error) {
switch contentEncoding {
case "", "identity":
return io.NopCloser(r), &nopCloserWrapper{w}, nil
case "gzip":
gzReader, err := gzip.NewReader(r)
if err != nil {
return nil, nil, trace.Wrap(err)
}
gzWriter := gzipPool.Get().(*gzip.Writer)
gzWriter.Reset(w)
return gzReader, &gzipWrapper{gzWriter}, nil
default:
return nil, nil, trace.BadParameter("unsupported Content-Encoding: %s", contentEncoding)
}
}
/*
* Teleport
* Copyright (C) 2024 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 proxy
import (
"bytes"
"context"
"net/http"
"strings"
jsonpatch "github.com/evanphx/json-patch"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime"
apimachinerytypes "k8s.io/apimachinery/pkg/types"
"k8s.io/apimachinery/pkg/util/strategicpatch"
"k8s.io/apimachinery/pkg/watch"
"k8s.io/client-go/kubernetes"
"k8s.io/client-go/rest"
"github.com/gravitational/teleport"
apidefaults "github.com/gravitational/teleport/api/defaults"
kubewaitingcontainerpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/kubewaitingcontainer/v1"
"github.com/gravitational/teleport/api/types/kubewaitingcontainer"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/utils"
)
// ephemeralContainers handles ephemeral container creation requests.
// If a user that is required to be moderated attempts to create an
// ephemeral container, the creation of that container will be delayed
// until the requirements for the moderated session are met.
func (f *Forwarder) ephemeralContainers(authCtx *authContext, w http.ResponseWriter, req *http.Request, p httprouter.Params) (resp any, err error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/ephemeralContainers",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("ephemeralContainers"),
semconv.RPCSystemKey.String("kube"),
),
)
req = req.WithContext(ctx)
defer span.End()
// If the user can start a session by themselves, proxy the ephemeral
// container creation request. Otherwise if the user requires
// moderation reply with fake data so kubectl will attempt to start
// a session with this ephemeral container. Then we will wait to
// create the ephemeral container until the requirements for the
// moderated session are met. If we wait here kubectl will timeout,
// so make it wait to establish a session instead.
canStart, err := f.canStartSessionAlone(authCtx)
if err != nil {
return nil, trace.Wrap(err)
}
if canStart {
return f.catchAll(authCtx, w, req)
}
sess, err := f.newClusterSession(req.Context(), *authCtx)
if err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(req.Context(), "Failed to create cluster session", "error", err)
return nil, trace.Wrap(err)
}
// sess.Close cancels the connection monitor context to release it sooner.
// When the server is under heavy load it can take a while to identify that
// the underlying connection is gone. This change prevents that and releases
// the resources as soon as we know the session is no longer active.
defer sess.close()
sess.upgradeToHTTP2 = true
sess.forwarder, err = f.makeSessionForwarder(sess)
if err != nil {
return nil, trace.Wrap(err)
}
if err := f.setupForwardingHeaders(sess, req, true /* withImpersonationHeaders */); err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(req.Context(), "Failed to set up forwarding headers", "error", err)
return nil, trace.Wrap(err)
}
if !sess.isLocalKubernetesCluster {
sess.forwarder.ServeHTTP(w, req)
return nil, nil
}
err = f.ephemeralContainersLocal(authCtx, sess, w, req)
return nil, trace.Wrap(err)
}
// ephemeralContainersLocal handles ephemeral container creation requests for
// users that require moderation.
func (f *Forwarder) ephemeralContainersLocal(authCtx *authContext, sess *clusterSession, w http.ResponseWriter, req *http.Request) (err error) {
// Fetch information on the requested pod and apply the patch
// so kubectl will think the ephemeral container has been created.
podPatch, err := utils.ReadAtMost(req.Body, teleport.MaxHTTPRequestSize)
if err != nil {
return trace.Wrap(err)
}
if err := req.Body.Close(); err != nil {
return trace.Wrap(err)
}
reqContentType := responsewriters.GetContentTypeHeader(req.Header)
// Remove "; charset=" if included in header.
if idx := strings.Index(reqContentType, ";"); idx > 0 {
reqContentType = reqContentType[:idx]
}
reqPatchType := apimachinerytypes.PatchType(reqContentType)
contentType, err := patchTypeToContentType(reqPatchType)
if err != nil {
return trace.Wrap(err)
}
encoder, decoder, err := newEncoderAndDecoderForContentType(
contentType,
newClientNegotiator(sess.codecFactory))
if err != nil {
return trace.Wrap(err, "failed to create encoder and decoder")
}
// Fetch the target pod using the user's impersonated identity so
// the Kubernetes API server enforces RBAC on `kubernetes_users` / `kubernetes_groups`.
pod, err := f.getPodForEphemeralPatch(
req.Context(),
authCtx,
req.Header,
authCtx.metaResource.requestedResource.namespace,
authCtx.metaResource.requestedResource.resourceName,
)
if err != nil {
return trace.Wrap(err)
}
patchedPod, ephemeralContName, err := f.mergeEphemeralPatchWithCurrentPod(
pod,
mergeEphemeralPatchWithCurrentPodConfig{
decoder: decoder,
encoder: encoder,
podPatch: podPatch,
patchType: reqPatchType,
},
)
if err != nil {
return trace.Wrap(err)
}
// Resolve the impersonation identity that the caller's headers select.
// Stored on the waiting container so the watch path can replay the same impersonation later,
// when the original request headers are gone.
kubeUser, kubeGroups, err := computeAndValidateImpersonatedPrincipals(
authCtx.kubeUsers, authCtx.kubeGroups, authCtx.User.GetName(), req.Header,
)
if err != nil {
return trace.Wrap(err)
}
if err := f.createWaitingContainer(req.Context(), ephemeralContName, authCtx, podPatch, reqPatchType, kubeUser, kubeGroups); err != nil {
return trace.Wrap(err)
}
responsewriters.SetContentTypeHeader(w, w.Header())
w.WriteHeader(http.StatusOK)
if err := encoder.Encode(patchedPod, w); err != nil {
return trace.Wrap(err)
}
f.emitAuditEvent(req, sess, http.StatusOK)
return trace.Wrap(err)
}
// mergeEphemeralPatchWithCurrentPodConfig is a configuration struct for
// mergeEphemeralPatchWithCurrentPod.
type mergeEphemeralPatchWithCurrentPodConfig struct {
decoder runtime.Decoder
encoder runtime.Encoder
podPatch []byte
patchType apimachinerytypes.PatchType
}
// mergeEphemeralPatchWithCurrentPod merges the provided patch with the
// given pod and returns the patched pod.
// The pod must have been fetched by the caller using a client that respects the requesting user's Kubernetes RBAC.
// The patch is expected to be a strategic merge patch that adds an ephemeral container to the pod.
func (f *Forwarder) mergeEphemeralPatchWithCurrentPod(
pod *corev1.Pod,
cfg mergeEphemeralPatchWithCurrentPodConfig,
) (*corev1.Pod, string, error) {
podSerializedBuf := &bytes.Buffer{}
if err := cfg.encoder.Encode(pod, podSerializedBuf); err != nil {
return nil, "", trace.Wrap(err)
}
patchedPod, ephemeralContName, err := patchPodWithDebugContainer(cfg.decoder, podSerializedBuf.Bytes(), cfg.podPatch, *pod, cfg.patchType)
if err != nil {
return nil, "", trace.Wrap(err)
}
return patchedPod, ephemeralContName, nil
}
// impersonationHeadersFromWaitingContainer materializes Impersonate-User/Impersonate-Group headers
// from the impersonation captured when the waiting container was created,
// so getPodForEphemeralPatch can re-apply the same choice on the watch path where no live request headers exist.
func impersonationHeadersFromWaitingContainer(waitingCont *kubewaitingcontainerpb.KubernetesWaitingContainer) http.Header {
headers := http.Header{}
if user := waitingCont.GetSpec().GetKubernetesUser(); user != "" {
headers.Set(ImpersonateUserHeader, user)
}
for _, group := range waitingCont.GetSpec().GetKubernetesGroups() {
headers.Add(ImpersonateGroupHeader, group)
}
return headers
}
// getPodForEphemeralPatch fetches the target pod using a Kubernetes client
// that impersonates the requesting user's `kubernetes_users` / `kubernetes_groups`.
// The Kubernetes API server therefore enforces RBAC on the user's mapped identity for this read, even though
// the surrounding moderated-session flow synthesizes the patch response locally without forwarding the PATCH itself.
func (f *Forwarder) getPodForEphemeralPatch(
ctx context.Context,
authCtx *authContext,
headers http.Header,
namespace, podName string,
) (*corev1.Pod, error) {
clientSet, _, err := f.impersonatedKubeClient(authCtx, headers)
if err != nil {
return nil, trace.Wrap(err)
}
pod, err := clientSet.CoreV1().
Pods(namespace).
Get(ctx, podName, metav1.GetOptions{})
if err != nil {
return nil, trace.Wrap(err)
}
return pod, nil
}
func (f *Forwarder) createWaitingContainer(ctx context.Context, ephemeralContName string, authCtx *authContext, podPatch []byte, patchType apimachinerytypes.PatchType, kubeUser string, kubeGroups []string) error {
waitingCont, err := kubewaitingcontainer.NewKubeWaitingContainer(
ephemeralContName,
kubewaitingcontainerpb.KubernetesWaitingContainerSpec_builder{
Username: authCtx.User.GetName(),
Cluster: authCtx.kubeClusterName,
Namespace: authCtx.metaResource.requestedResource.namespace,
PodName: authCtx.metaResource.requestedResource.resourceName,
ContainerName: ephemeralContName,
Patch: podPatch,
PatchType: string(patchType),
KubernetesUser: kubeUser,
KubernetesGroups: kubeGroups,
}.Build())
if err != nil {
return trace.Wrap(err)
}
_, err = f.cfg.AuthClient.CreateKubernetesWaitingContainer(ctx, waitingCont)
return trace.Wrap(err)
}
// impersonatedKubeClient returns a Kubernetes client that is impersonating
// the identity in the provided authCtx.
func (f *Forwarder) impersonatedKubeClient(authCtx *authContext, headers http.Header) (*kubernetes.Clientset, *kubeDetails, error) {
details, err := f.findKubeDetailsByClusterName(authCtx.kubeClusterName)
if err != nil {
return nil, nil, trace.NotFound("kubernetes cluster %q not found", authCtx.kubeClusterName)
}
kubeUser, kubeGroups, err := computeAndValidateImpersonatedPrincipals(authCtx.kubeUsers, authCtx.kubeGroups, authCtx.User.GetName(), headers)
if err != nil {
return nil, nil, trace.Wrap(err)
}
// Clone the shared cached rest config before setting impersonation to avoid
// racing with concurrent requests to the same cluster.
restConfig := *details.getKubeRestConfig()
restConfig.Impersonate = rest.ImpersonationConfig{
UserName: kubeUser,
Groups: kubeGroups,
}
clientSet, err := kubernetes.NewForConfig(&restConfig)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return clientSet, details, nil
}
// patchPodWithDebugContainer adds an ephemeral container to the provided spec of pod and
// returns the patched result.
func patchPodWithDebugContainer(decoder runtime.Decoder, podJson, podPatch []byte, pod corev1.Pod, patchType apimachinerytypes.PatchType) (*corev1.Pod, string, error) {
patchResult, err := patchPod(podJson, podPatch, pod, patchType)
if err != nil {
return nil, "", trace.Wrap(err)
}
gvk := corev1.SchemeGroupVersion.WithKind("Pod")
decodedObj, _, err := decoder.Decode(patchResult, &gvk, nil)
if err != nil {
return nil, "", trace.Wrap(err)
}
decodedObj.GetObjectKind().SetGroupVersionKind(gvk)
patchedPod, ok := decodedObj.(*corev1.Pod)
if !ok {
return nil, "", trace.CompareFailed("expected *corev1.Pod, got %T", decodedObj)
}
// Determine which ephemeral containers the patch added relative to the
// pre-patch pod. A strategic merge patch can append more than one entry,
// so inspecting only the final element is not enough: every added
// container must be accounted for and validated.
existing := make(map[string]struct{}, len(pod.Spec.EphemeralContainers))
for _, c := range pod.Spec.EphemeralContainers {
existing[c.Name] = struct{}{}
}
var added []corev1.EphemeralContainer
for _, c := range patchedPod.Spec.EphemeralContainers {
if _, ok := existing[c.Name]; !ok {
added = append(added, c)
}
}
if len(added) != 1 {
return nil, "", trace.AccessDenied("exactly one ephemeral container may be added per request, got %d", len(added))
}
ephemeralCont := added[0]
if !ephemeralCont.TTY {
return nil, "", trace.AccessDenied("only interactive ephemeral containers are supported")
}
// Add the container to the status so kubectl will think it has started.
patchedPod.Status.EphemeralContainerStatuses = append(
pod.Status.EphemeralContainerStatuses,
corev1.ContainerStatus{
Name: ephemeralCont.Name,
State: corev1.ContainerState{
Running: &corev1.ContainerStateRunning{
StartedAt: metav1.Now(),
},
},
Ready: true,
},
)
return patchedPod, ephemeralCont.Name, nil
}
// pushPodEvent writes a fake event that shows that an ephemeral container
// started running on a given pod. This is so kubectl will attempt to start
// a session which can be safely waiting on until the moderated session
// is approved.
func (f *Forwarder) getPatchedPodEvent(ctx context.Context, sess *clusterSession, waitingCont *kubewaitingcontainerpb.KubernetesWaitingContainer) (*watch.Event, error) {
contentType, err := patchTypeToContentType(apimachinerytypes.PatchType(waitingCont.GetSpec().GetPatchType()))
if err != nil {
return nil, trace.Wrap(err)
}
encoder, decoder, err := newEncoderAndDecoderForContentType(
contentType,
newClientNegotiator(sess.codecFactory),
)
if err != nil {
return nil, trace.Wrap(err, "failed to create encoder and decoder")
}
pod, err := f.getPodForEphemeralPatch(
ctx,
&sess.authContext,
impersonationHeadersFromWaitingContainer(waitingCont),
waitingCont.GetSpec().GetNamespace(),
waitingCont.GetSpec().GetPodName(),
)
if err != nil {
return nil, trace.Wrap(err)
}
patchedPod, _, err := f.mergeEphemeralPatchWithCurrentPod(
pod,
mergeEphemeralPatchWithCurrentPodConfig{
decoder: decoder,
encoder: encoder,
podPatch: waitingCont.GetSpec().GetPatch(),
patchType: apimachinerytypes.PatchType(waitingCont.GetSpec().GetPatchType()),
},
)
if err != nil {
return nil, trace.Wrap(err)
}
return &watch.Event{
Type: watch.Modified,
Object: patchedPod,
}, nil
}
// getUserEphemeralContainersForPod returns a list of ephemeral containers
// created by the username and are waiting to be created for a given pod.
func (f *Forwarder) getUserEphemeralContainersForPod(ctx context.Context, username, kubeCluster, namespace, pod string) ([]*kubewaitingcontainerpb.KubernetesWaitingContainer, error) {
if f.cfg.GetScope() != "" {
// If the kube forwarder is scoped then moderated sessions are not supported and access to
// KindKubernetesWaitingContainer will be denied. We need to return without error to prevent
// interactive exec from failing
return nil, nil
}
var (
list []*kubewaitingcontainerpb.KubernetesWaitingContainer
startPage string
)
for {
waitingContainers, nextPage, err := f.cfg.CachingAuthClient.ListKubernetesWaitingContainers(ctx, apidefaults.DefaultChunkSize, startPage)
if err != nil {
return nil, trace.Wrap(err)
}
for _, cont := range waitingContainers {
if cont.GetSpec().GetUsername() != username ||
cont.GetSpec().GetCluster() != kubeCluster ||
cont.GetSpec().GetNamespace() != namespace ||
cont.GetSpec().GetPodName() != pod {
continue
}
list = append(list, cont)
}
if nextPage == "" {
break
}
startPage = nextPage
}
return list, nil
}
func getEphemeralContainerStatusByName(pod *corev1.Pod, containerName string) *corev1.ContainerStatus {
for _, status := range pod.Status.EphemeralContainerStatuses {
if status.Name == containerName {
return &status
}
}
return nil
}
func patchTypeToContentType(reqPatchType apimachinerytypes.PatchType) (string, error) {
var contentType string
switch reqPatchType {
case apimachinerytypes.JSONPatchType,
apimachinerytypes.MergePatchType,
apimachinerytypes.StrategicMergePatchType:
contentType = responsewriters.JSONContentType
case apimachinerytypes.ApplyPatchType:
contentType = responsewriters.YAMLContentType
default:
return "", trace.BadParameter("unsupported content type %q", reqPatchType)
}
return contentType, nil
}
// patchPod applies the provided patch to the pod and returns the patched pod data.
// The patch type is used to determine how the patch should be applied.
func patchPod(podData, patchData []byte, pod corev1.Pod, pt apimachinerytypes.PatchType) ([]byte, error) {
switch pt {
case apimachinerytypes.JSONPatchType:
patchObj, err := jsonpatch.DecodePatch(patchData)
if err != nil {
return nil, trace.Wrap(err)
}
patchedObj, err := patchObj.Apply(podData)
return patchedObj, trace.Wrap(err)
case apimachinerytypes.MergePatchType:
patchedObj, err := jsonpatch.MergePatch(podData, patchData)
return patchedObj, trace.Wrap(err)
case apimachinerytypes.StrategicMergePatchType:
patchedObj, err := strategicpatch.StrategicMergePatch(podData, patchData, pod)
return patchedObj, trace.Wrap(err)
default:
return nil, trace.BadParameter("unsupported patch type %q", pt)
}
}
/*
* Teleport
* Copyright (C) 2024 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 proxy
import (
"net/http"
kubeerrors "k8s.io/apimachinery/pkg/api/errors"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
)
const (
// kubernetesSessionTerminatedByUser is the message that is sent to the
// client when the session is terminated by the moderator.
kubernetesSessionTerminatedByModerator = "Session terminated by moderator."
sessionTerminatedByModeratorReason = metav1.StatusReason("SessionTerminatedByModerator")
)
var sessionTerminatedByModeratorErr = &kubeerrors.StatusError{
ErrStatus: metav1.Status{
Status: metav1.StatusFailure,
Code: http.StatusUnauthorized,
Reason: sessionTerminatedByModeratorReason,
Message: kubernetesSessionTerminatedByModerator,
Details: &metav1.StatusDetails{
Causes: []metav1.StatusCause{
{
Type: metav1.CauseTypeForbidden,
Message: kubernetesSessionTerminatedByModerator,
},
},
},
},
}
// isSessionTerminatedError returns true if the error is a session terminated error.
// This is required because StreamWithContext wraps the error into a new error string
// and we lose the type information to forward the error to the client.
func isSessionTerminatedError(err error) bool {
if err == nil {
return false
}
// This check is required because the error is wrapped into a new error string
// by StreamWithContext and we lose the type information.
return err.Error() == kubernetesSessionTerminatedByModerator
}
/*
* Teleport
* Copyright (C) 2026 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 proxy
import (
"regexp"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// fastMatcher is a precompiled per-request RBAC matcher that reduces per-item matching overhead.
// Instead of calling matchKubernetesResource per item (which does cache lookups, rule iteration, per-field matching),
// the fast matcher resolves constant fields at compile time and only checks per-item fields during matching.
//
// Of the five rule fields:
// - Kind: exact-match or wildcard only, constant per request, resolved at compile time.
// - Verb: exact-match or wildcard only, constant per request, resolved at compile time.
// - APIGroup: supports regex/glob, constant per request, resolved at compile time.
// - Namespace: supports regex/glob, varies per item, checked per item.
// - Name: supports regex/glob, varies per item, checked per item.
//
// The fast matcher handles the common "default" case in KubeResourceMatchesRegex.
// It cannot handle namespace special cases (when the requested kind is "namespaces").
type fastMatcher struct {
allowRules []compiledMatchRule
denyRules []compiledMatchRule
}
// compiledMatchRule is a single RBAC rule with pre-compiled name and namespace matchers.
// Kind, verb, and apiGroup have already been resolved during rule filtering.
type compiledMatchRule struct {
name fieldMatcher
namespace fieldMatcher
// requiresNamespace is true when the original rule had a non-empty, non-wildcard namespace pattern.
// When true, resources with an empty namespace cannot match this rule.
requiresNamespace bool
}
// fieldMatcher matches a string either by exact literal comparison or by compiled regex.
type fieldMatcher struct {
literal string // set for exact match
re *regexp.Regexp // set for pattern match
}
func (f fieldMatcher) match(s string) bool {
if f.re == nil {
return f.literal == s
}
return f.re.MatchString(s)
}
// newFastMatcher filters rules by request parameters and compiles a fast matcher.
func newFastMatcher(mr metaResource, allowed, denied []types.KubernetesResource) (*fastMatcher, error) {
// Local cache for this compilation pass.
// Many rules share the same patterns, so this avoids compiling the same expression multiple times.
cache := make(map[string]*regexp.Regexp)
filteredAllowed, err := filterRules(mr, allowed, cache)
if err != nil {
return nil, trace.Wrap(err)
}
filteredDenied, err := filterRules(mr, denied, cache)
if err != nil {
return nil, trace.Wrap(err)
}
allowRules, err := compileMatchRules(filteredAllowed, cache)
if err != nil {
return nil, trace.Wrap(err)
}
denyRules, err := compileMatchRules(filteredDenied, cache)
if err != nil {
return nil, trace.Wrap(err)
}
return &fastMatcher{
allowRules: allowRules,
denyRules: denyRules,
}, nil
}
// filterRules returns the subset of rules that match the given request parameters.
func filterRules(mr metaResource, rules []types.KubernetesResource, cache map[string]*regexp.Regexp) ([]types.KubernetesResource, error) {
filtered := make([]types.KubernetesResource, 0, len(rules))
for _, r := range rules {
if !kindAllowed(r.Kind, mr.requestedResource.resourceKind) {
continue
}
if !verbAllowed(r.Verbs, mr.verb) {
continue
}
if !namespaceAllowed(r.Namespace, mr.requestedResource.namespace) {
continue
}
match, err := apiGroupMatches(r.APIGroup, mr.requestedResource.apiGroup, cache)
if err != nil {
return nil, trace.Wrap(err)
}
if !match {
continue
}
filtered = append(filtered, r)
}
return filtered, nil
}
func kindAllowed(ruleKind, requestedKind string) bool {
return ruleKind == types.Wildcard || ruleKind == requestedKind
}
func verbAllowed(allowedVerbs []string, verb string) bool {
return utils.IsVerbAllowed(allowedVerbs, verb)
}
func namespaceAllowed(ruleNamespace, requestedNamespace string) bool {
if requestedNamespace == "" {
// Cluster-wide request: all rules pass since items may come from any namespace.
return true
}
// Empty rule namespace targets cluster-wide resources.
// Keep it during pre-filtering; the compiled matcher only matches items with empty namespace.
if ruleNamespace == "" {
return true
}
// Pattern namespaces are kept since they need compilation.
if isGlobOrRegexp(ruleNamespace) {
return true
}
return ruleNamespace == requestedNamespace
}
func apiGroupMatches(ruleAPIGroup, requestedAPIGroup string, cache map[string]*regexp.Regexp) (bool, error) {
if !isGlobOrRegexp(ruleAPIGroup) {
return ruleAPIGroup == requestedAPIGroup, nil
}
if re, ok := cache[ruleAPIGroup]; ok {
return re.MatchString(requestedAPIGroup), nil
}
re, err := utils.CompileExpression(ruleAPIGroup)
if err != nil {
return false, trace.Wrap(err)
}
cache[ruleAPIGroup] = re
return re.MatchString(requestedAPIGroup), nil
}
func isGlobOrRegexp(expr string) bool {
return strings.Contains(expr, "*") || utils.IsRegexp(expr)
}
func compileMatchRules(resources []types.KubernetesResource, cache map[string]*regexp.Regexp) ([]compiledMatchRule, error) {
rules := make([]compiledMatchRule, 0, len(resources))
for _, r := range resources {
nameM, err := compileFieldMatcher(r.Name, cache)
if err != nil {
return nil, trace.Wrap(err)
}
nsM, err := compileFieldMatcher(r.Namespace, cache)
if err != nil {
return nil, trace.Wrap(err)
}
rules = append(rules, compiledMatchRule{
name: nameM,
namespace: nsM,
requiresNamespace: r.Namespace != "" && r.Namespace != types.Wildcard,
})
}
return rules, nil
}
// compileFieldMatcher returns a fieldMatcher for the given expression.
// Literal expressions (no wildcards or regex) use direct string comparison.
// Patterns are compiled to regex, with results cached across rules.
func compileFieldMatcher(expression string, cache map[string]*regexp.Regexp) (fieldMatcher, error) {
if !isGlobOrRegexp(expression) {
return fieldMatcher{literal: expression}, nil
}
if re, ok := cache[expression]; ok {
return fieldMatcher{re: re}, nil
}
re, err := utils.CompileExpression(expression)
if err != nil {
return fieldMatcher{}, trace.Wrap(err)
}
cache[expression] = re
return fieldMatcher{re: re}, nil
}
// Match checks if a resource with the given name and namespace is allowed by the precompiled RBAC rules.
func (m *fastMatcher) Match(name, namespace string) (bool, error) {
for i := range m.denyRules {
if m.denyRules[i].matches(name, namespace) {
return false, nil
}
}
for i := range m.allowRules {
if m.allowRules[i].matches(name, namespace) {
return true, nil
}
}
return false, nil
}
// matches checks whether a single compiled rule matches the given fields.
// This mirrors the "default" case in KubeResourceMatchesRegex.
// Kind, verb, and apiGroup are already resolved during rule filtering.
func (r *compiledMatchRule) matches(name, namespace string) bool {
if !r.name.match(name) {
return false
}
if r.requiresNamespace && namespace == "" {
return false
}
return r.namespace.match(namespace)
}
/*
* Teleport
* Copyright (C) 2026 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 proxy
import (
"context"
"io"
"log/slog"
"maps"
"net/http"
"strings"
"sync"
"github.com/gravitational/trace"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/kube/proxy/streamfilter"
)
// filteringResponseWriter is an http.ResponseWriter that intercepts the upstream response,
// inspects headers to decide between streaming and buffered filtering, and routes the body accordingly.
type filteringResponseWriter struct {
target http.ResponseWriter
headers http.Header
status int
body io.Writer
once sync.Once
matcher resourceMatcher
filterWrapper responsewriters.FilterWrapper
log *slog.Logger
ctx context.Context
tracer oteltrace.Tracer
kubeServiceType string
pipeWriter *io.PipeWriter
memBuffer *responsewriters.MemoryResponseWriter
filterErr error
filterDone chan struct{} // closed when the streaming goroutine finishes
streaming bool
}
func newFilteringResponseWriter(
target http.ResponseWriter,
matcher resourceMatcher,
filterWrapper responsewriters.FilterWrapper,
log *slog.Logger,
ctx context.Context,
tracer oteltrace.Tracer,
kubeServiceType string,
) *filteringResponseWriter {
return &filteringResponseWriter{
target: target,
headers: make(http.Header),
matcher: matcher,
filterWrapper: filterWrapper,
log: log,
ctx: ctx,
tracer: tracer,
kubeServiceType: kubeServiceType,
}
}
func (fw *filteringResponseWriter) Header() http.Header {
return fw.headers
}
func (fw *filteringResponseWriter) WriteHeader(statusCode int) {
fw.once.Do(func() {
fw.status = statusCode
if statusCode == http.StatusOK {
if fw.tryStreaming() {
return
}
}
fw.memBuffer = responsewriters.NewMemoryResponseWriter()
maps.Copy(fw.memBuffer.Header(), fw.headers)
fw.memBuffer.WriteHeader(statusCode)
fw.body = fw.memBuffer.Buffer()
})
}
// tryStreaming attempts to set up the streaming filter path.
// Returns true if streaming was activated, false to fall back to buffered.
func (fw *filteringResponseWriter) tryStreaming() bool {
contentType := responsewriters.GetContentTypeHeader(fw.headers)
if !strings.Contains(contentType, "application/json") {
return false
}
sf := streamfilter.NewJSONFilter(fw.matcher, fw.log)
contentEncoding := fw.headers.Get("Content-Encoding")
if contentEncoding != "" && contentEncoding != "identity" && contentEncoding != "gzip" {
fw.log.WarnContext(fw.ctx, "Unexpected Content-Encoding, falling back to buffered filter", "content_encoding", contentEncoding)
return false
}
maps.Copy(fw.target.Header(), fw.headers)
fw.target.Header().Del("Content-Length")
fw.target.WriteHeader(fw.status)
pr, pw := io.Pipe()
fw.pipeWriter = pw
fw.streaming = true
fw.filterDone = make(chan struct{})
fw.body = pw
// wrapContentEncoding is called inside the goroutine because gzip.NewReader
// reads the gzip header from the pipe, which blocks until Write provides data.
// Calling it here in WriteHeader would deadlock.
go func() {
defer close(fw.filterDone)
// Close the read end of the pipe when done to unblock any pending
// pipeWriter.Write in ServeHTTP (e.g. if the filter exits early
// due to a client disconnect).
defer pr.Close()
src, dst, err := wrapContentEncoding(pr, fw.target, contentEncoding)
if err != nil {
fw.filterErr = trace.ConnectionProblem(err, "failed to initialize content encoding wrapper for %q", contentEncoding)
fw.log.ErrorContext(fw.ctx, "Streaming filter content encoding setup failed", "error", fw.filterErr)
return
}
filterErr := sf.Filter(src, dst)
fw.filterErr = trace.NewAggregate(filterErr, dst.Close(), src.Close())
if fw.filterErr != nil {
fw.log.ErrorContext(fw.ctx, "Streaming filter failed mid-write, client received truncated response", "error", fw.filterErr)
}
}()
return true
}
func (fw *filteringResponseWriter) Write(b []byte) (int, error) {
fw.WriteHeader(http.StatusOK)
return fw.body.Write(b)
}
// Finish completes filtering and returns the status code and any error.
func (fw *filteringResponseWriter) Finish() (int, error) {
if fw.status == 0 {
return http.StatusBadGateway, trace.ConnectionProblem(nil, "upstream closed without response")
}
if fw.streaming {
fw.pipeWriter.Close()
<-fw.filterDone
return fw.status, fw.filterErr
}
_, filterSpan := fw.tracer.Start(fw.ctx, "kube.Forwarder/listResourcesList/filterBuffer",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(fw.kubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
err := filterBuffer(fw.filterWrapper, fw.memBuffer)
filterSpan.End()
if err != nil {
return fw.memBuffer.Status(), trace.Wrap(err)
}
err = fw.memBuffer.CopyInto(fw.target)
return fw.memBuffer.Status(), trace.Wrap(err)
}
/*
* 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 proxy
import (
"cmp"
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net"
"net/http"
"net/url"
"slices"
"strconv"
"strings"
"sync"
"time"
"github.com/google/uuid"
gwebsocket "github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/julienschmidt/httprouter"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
"golang.org/x/net/http/httpguts"
kubeerrors "k8s.io/apimachinery/pkg/api/errors"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/apimachinery/pkg/runtime/serializer"
utilnet "k8s.io/apimachinery/pkg/util/net"
"k8s.io/client-go/rest"
"k8s.io/client-go/tools/portforward"
"k8s.io/client-go/tools/remotecommand"
"k8s.io/client-go/transport/spdy"
kwebsocket "k8s.io/client-go/transport/websocket"
kubeexec "k8s.io/client-go/util/exec"
"k8s.io/streaming/pkg/httpstream"
httpstreamspdy "k8s.io/streaming/pkg/httpstream/spdy"
"k8s.io/streaming/pkg/httpstream/wsstream"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/api/observability/tracing"
tracehttp "github.com/gravitational/teleport/api/observability/tracing/http"
"github.com/gravitational/teleport/api/types"
apievents "github.com/gravitational/teleport/api/types/events"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/entitlements"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/auth/moderation"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/httplib/reverseproxy"
"github.com/gravitational/teleport/lib/kube/internal"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/kube/proxy/streamproto"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/multiplexer"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/scopes/pinning"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/set"
)
// KubeServiceType specifies a Teleport service type which can forward Kubernetes requests
type KubeServiceType = string
const (
// KubeService is a Teleport kubernetes_service. A KubeService always forwards
// requests directly to a Kubernetes endpoint.
KubeService = "kube_service"
// ProxyService is a Teleport proxy_service with kube_listen_addr/
// kube_public_addr enabled. A ProxyService always forwards requests to a
// Teleport KubeService or LegacyProxyService.
ProxyService = "kube_proxy"
// LegacyProxyService is a Teleport proxy_service with the kubernetes section
// enabled. A LegacyProxyService can forward requests directly to a Kubernetes
// endpoint, or to another Teleport LegacyProxyService or KubeService.
LegacyProxyService = "legacy_proxy"
)
// ForwarderConfig specifies configuration for proxy forwarder
type ForwarderConfig struct {
// ReverseTunnelSrv is the teleport reverse tunnel server
ReverseTunnelSrv reversetunnelclient.Server
// ClusterName is a local cluster name
ClusterName string
// Keygen points to a key generator implementation
Keygen sshca.Authority
// ScopedAuthz authenticates user
ScopedAuthz authz.ScopedAuthorizer
// AuthClient is a auth server client.
AuthClient authclient.ClientI
// CachingAuthClient is a caching auth server client for read-only access.
CachingAuthClient authclient.ReadKubernetesAccessPoint
// Emitter is used to emit audit events
Emitter apievents.Emitter
// DataDir is a data dir to store logs
DataDir string
// Namespace is a namespace of the proxy server (not a K8s namespace)
Namespace string
// HostID is a unique ID of a proxy server
HostID string
// ClusterOverride if set, routes all requests
// to the cluster name, used in tests
ClusterOverride string
// Context passes the optional external context
// passing global close to all forwarder operations
Context context.Context
// KubeconfigPath is a path to kubernetes configuration
KubeconfigPath string
// KubeServiceType specifies which Teleport service type this forwarder is for
KubeServiceType KubeServiceType
// KubeClusterName is the name of the kubernetes cluster that this
// forwarder handles.
KubeClusterName string
// Clock is a server clock, could be overridden in tests
Clock clockwork.Clock
// ConnPingPeriod is a period for sending ping messages on the incoming
// connection.
ConnPingPeriod time.Duration
// Component name to include in log output.
Component string
// LockWatcher is a lock watcher.
LockWatcher *services.LockWatcher
// CheckImpersonationPermissions is an optional override of the default
// impersonation permissions check, for use in testing
CheckImpersonationPermissions servicecfg.ImpersonationPermissionsChecker
// PublicAddr is the address that can be used to reach the kube cluster
PublicAddr string
// PROXYSigner is used to sign PROXY headers for securely propagating client IP address
PROXYSigner multiplexer.PROXYHeaderSigner
// log is the logger function
log *slog.Logger
// TracerProvider is used to create tracers capable
// of starting spans.
TracerProvider oteltrace.TracerProvider
// Tracer is used to start spans.
tracer oteltrace.Tracer
// GetConnTLSCertificate returns the TLS client certificate to use when
// connecting to the upstream Teleport proxy or Kubernetes service when
// forwarding requests using the forward identity (i.e. proxy impersonating
// a user) method. Paired with GetConnTLSRoots and ConnTLSCipherSuites to
// generate the correct [*tls.Config] on demand.
GetConnTLSCertificate utils.GetCertificateFunc
// GetConnTLSRoots returns the [*x509.CertPool] used to validate TLS
// connections to the upstream Teleport proxy or Kubernetes service.
GetConnTLSRoots utils.GetRootsFunc
// ConnTLSCipherSuites optionally contains a list of TLS ciphersuites to use
// when connecting to the upstream Teleport Proxy or Kubernetes service.
ConnTLSCipherSuites []uint16
// ClusterFeaturesGetter is a function that returns the Teleport cluster licensed features.
// It is used to determine if the cluster is licensed for Kubernetes usage.
ClusterFeatures ClusterFeaturesGetter
// Scope is the scope the forwarder is pinned to if a full scope pin is not present.
Scope string
// ScopePin is the scope and scoped role assignments the forwarder is pinned to.
ScopePin *scopesv1.Pin
}
// ClusterFeaturesGetter is a function that returns the Teleport cluster licensed features.
type ClusterFeaturesGetter func() proto.Features
func (f ClusterFeaturesGetter) GetEntitlement(e entitlements.EntitlementKind) modules.EntitlementInfo {
al, ok := f().Entitlements[string(e)]
if !ok {
return modules.EntitlementInfo{}
}
return modules.EntitlementInfo{
Enabled: al.Enabled,
Limit: al.Limit,
}
}
// CheckAndSetDefaults checks and sets default values
func (f *ForwarderConfig) CheckAndSetDefaults() error {
if f.AuthClient == nil {
return trace.BadParameter("missing parameter AuthClient")
}
if f.CachingAuthClient == nil {
return trace.BadParameter("missing parameter CachingAuthClient")
}
if f.ScopedAuthz == nil {
return trace.BadParameter("missing parameter ScopedAuthz")
}
if f.LockWatcher == nil {
return trace.BadParameter("missing parameter LockWatcher")
}
if f.Emitter == nil {
return trace.BadParameter("missing parameter Emitter")
}
if f.ClusterName == "" {
return trace.BadParameter("missing parameter ClusterName")
}
if f.Keygen == nil {
return trace.BadParameter("missing parameter Keygen")
}
if f.DataDir == "" {
return trace.BadParameter("missing parameter DataDir")
}
if f.HostID == "" {
return trace.BadParameter("missing parameter ServerID")
}
if f.ClusterFeatures == nil {
return trace.BadParameter("missing parameter ClusterFeatures")
}
if f.KubeServiceType != KubeService && f.PROXYSigner == nil {
return trace.BadParameter("missing parameter PROXYSigner")
}
if f.Namespace == "" {
f.Namespace = apidefaults.Namespace
}
if f.Context == nil {
f.Context = context.TODO()
}
if f.Clock == nil {
f.Clock = clockwork.NewRealClock()
}
if f.ConnPingPeriod == 0 {
f.ConnPingPeriod = defaults.HighResPollingPeriod
}
if f.Component == "" {
f.Component = "kube_forwarder"
}
if f.CheckImpersonationPermissions == nil {
f.CheckImpersonationPermissions = checkImpersonationPermissions
}
if f.TracerProvider == nil {
f.TracerProvider = tracing.DefaultProvider()
}
f.tracer = f.TracerProvider.Tracer("kube")
switch f.KubeServiceType {
case KubeService:
case ProxyService, LegacyProxyService:
if f.GetConnTLSCertificate == nil {
return trace.BadParameter("missing parameter GetConnTLSCertificate")
}
if f.GetConnTLSRoots == nil {
return trace.BadParameter("missing parameter GetConnTLSRoots")
}
default:
return trace.BadParameter("unknown value for KubeServiceType")
}
if f.KubeClusterName == "" && f.KubeconfigPath == "" && f.KubeServiceType == LegacyProxyService {
// Running without a kubeconfig and explicit k8s cluster name. Use
// teleport cluster name instead, to ask kubeutils.GetKubeConfig to
// attempt loading the in-cluster credentials.
f.KubeClusterName = f.ClusterName
}
if f.log == nil {
f.log = slog.Default()
}
if f.ScopePin != nil {
if err := pinning.WeakValidate(f.ScopePin); err != nil {
return trace.Wrap(err)
}
}
if f.Scope != "" {
if err := scopes.WeakValidate(f.Scope); err != nil {
return trace.Wrap(err)
}
}
if f.ScopePin.GetScope() != "" && f.Scope != "" {
return trace.BadParameter("either a scope pin or a bare scope must be set for a scoped kube forwarder, not both")
}
return nil
}
// GetScope returns the scope the forwarder is pinned to whether it's a bare scope or a scope pin.
func (f *ForwarderConfig) GetScope() string {
return cmp.Or(f.ScopePin.GetScope(), f.Scope)
}
// transportCacheTTL is the TTL for the transport cache.
const transportCacheTTL = 5 * time.Hour
// NewForwarder returns new instance of Kubernetes request
// forwarding proxy.
func NewForwarder(cfg ForwarderConfig) (*Forwarder, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
// TODO (tigrato): remove this once we have a better way to handle
// deleting expired entried clusters and kube_servers entries.
// In the meantime, we need to make sure that the cache is cleaned
// from time to time.
transportClients, err := utils.NewFnCache(utils.FnCacheConfig{
TTL: transportCacheTTL,
Clock: cfg.Clock,
})
if err != nil {
return nil, trace.Wrap(err)
}
closeCtx, close := context.WithCancel(cfg.Context)
fwd := &Forwarder{
log: cfg.log,
cfg: cfg,
activeRequests: make(map[string]context.Context),
ctx: closeCtx,
close: close,
sessions: make(map[uuid.UUID]*session),
upgrader: gwebsocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
},
clusterDetails: make(map[string]*kubeDetails),
cachedTransport: transportClients,
}
router := httprouter.New()
router.UseRawPath = true
router.GET("/version", fwd.withAuth(
func(ctx *authContext, w http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
// Forward version requests to the cluster.
return fwd.catchAll(ctx, w, r)
},
withCustomErrFormatter(fwd.writeResponseErrorToBody),
))
router.POST("/api/:ver/namespaces/:podNamespace/pods/:podName/exec", fwd.withAuth(fwd.exec))
router.GET("/api/:ver/namespaces/:podNamespace/pods/:podName/exec", fwd.withAuth(fwd.exec))
router.POST("/api/:ver/namespaces/:podNamespace/pods/:podName/attach", fwd.withAuth(fwd.exec))
router.GET("/api/:ver/namespaces/:podNamespace/pods/:podName/attach", fwd.withAuth(fwd.exec))
router.POST("/api/:ver/namespaces/:podNamespace/pods/:podName/portforward", fwd.withAuth(fwd.portForward))
router.GET("/api/:ver/namespaces/:podNamespace/pods/:podName/portforward", fwd.withAuth(fwd.portForward))
router.POST("/apis/authorization.k8s.io/:ver/selfsubjectaccessreviews", fwd.withAuth(fwd.selfSubjectAccessReviews))
router.PATCH("/api/:ver/namespaces/:podNamespace/pods/:podName/ephemeralcontainers", fwd.withAuth(fwd.ephemeralContainers))
router.PUT("/api/:ver/namespaces/:podNamespace/pods/:podName/ephemeralcontainers", fwd.withAuth(fwd.ephemeralContainers))
router.GET("/api/:ver/teleport/join/:session", fwd.withAuthPassthrough(fwd.join))
for _, method := range allHTTPMethods() {
router.Handle(method, "/v1/teleport/:base64Cluster/:base64KubeCluster/*path", fwd.singleCertHandler())
}
router.NotFound = fwd.withAuthStd(fwd.catchAll)
fwd.router = instrumentHTTPHandler(fwd.cfg.KubeServiceType, router)
if cfg.ClusterOverride != "" {
fwd.log.DebugContext(closeCtx, "Cluster override is set, forwarder will send all requests to remote cluster", "cluster_override", cfg.ClusterOverride)
}
if len(cfg.KubeClusterName) > 0 || len(cfg.KubeconfigPath) > 0 || cfg.KubeServiceType != KubeService {
if err := fwd.getKubeDetails(cfg.Context); err != nil {
return nil, trace.Wrap(err)
}
}
return fwd, nil
}
// Forwarder intercepts kubernetes requests, acting as Kubernetes API proxy.
// it blindly forwards most of the requests on HTTPS protocol layer,
// however some requests like exec sessions it intercepts and records.
type Forwarder struct {
mu sync.Mutex
log *slog.Logger
router http.Handler
cfg ForwarderConfig
// activeRequests is a map used to serialize active CSR requests to the auth server
activeRequests map[string]context.Context
// close is a close function
close context.CancelFunc
// ctx is a global context signaling exit
ctx context.Context
// clusterDetails contain kubernetes credentials for multiple clusters.
// map key is cluster name.
clusterDetails map[string]*kubeDetails
rwMutexDetails sync.RWMutex
// sessions tracks in-flight sessions
sessions map[uuid.UUID]*session
// upgrades connections to websockets
upgrader gwebsocket.Upgrader
// getKubernetesServersForKubeCluster is a function that returns a list of
// kubernetes servers for a given kube cluster but uses different methods
// depending on the service type.
// For example, if the service type is KubeService, it will use the
// local kubernetes clusters. If the service type is Proxy, it will
// use the heartbeat clusters.
getKubernetesServersForKubeCluster getKubeServersByNameFunc
// cachedTransport is a cache of cachedTransportEntry objects used to
// connect to Teleport services.
// TODO(tigrato): Implement a cache eviction policy using watchers.
cachedTransport *utils.FnCache
}
// cachedTransportEntry is a cached transport entry used to connect to
// Teleport services. It contains a cached http.RoundTripper and a cached
// tls.Config.
type cachedTransportEntry struct {
transport http.RoundTripper
tlsConfig *tls.Config
}
// getKubeServersByNameFunc is a function that returns a list of
// kubernetes servers for a given kube cluster.
type getKubeServersByNameFunc = func(ctx context.Context, name string) ([]types.KubeServer, error)
// Close signals close to all outstanding or background operations
// to complete
func (f *Forwarder) Close() error {
f.close()
return nil
}
func (f *Forwarder) ServeHTTP(rw http.ResponseWriter, r *http.Request) {
f.router.ServeHTTP(rw, r)
}
// authContext is a context of authenticated user,
// contains information about user, target cluster and authenticated groups
type authContext struct {
*authz.ScopedContext
// checker will be set once we know which access checker has been used to grant access
checker *services.ScopedAccessChecker
accessState services.AccessState
kubeGroups map[string]struct{}
kubeUsers map[string]struct{}
kubeClusterLabels map[string]string
kubeClusterName string
teleportCluster teleportClusterClient
recordingConfig types.SessionRecordingConfig
// clientIdleTimeout sets information on client idle timeout
clientIdleTimeout time.Duration
// clientIdleTimeoutMessage is the message to be displayed to the user
// when the client idle timeout is reached
clientIdleTimeoutMessage string
// disconnectExpiredCert if set, controls the time when the connection
// should be disconnected because the client cert expires
disconnectExpiredCert time.Time
// certExpires is the client certificate expiration timestamp.
certExpires time.Time
// sessionTTL specifies the duration of the user's session
sessionTTL time.Duration
// kubeCluster is the Kubernetes cluster the request is targeted to.
// It's only available after authorization layer.
kubeCluster types.KubeCluster
// metaResource holds the resource data:
// - the requested resource
// - the looked up resource definition, including the flag to know if it is namespaced
// - the verb used to access the resource
metaResource metaResource
// kubeServers are the registered agents for the kubernetes cluster the request
// is targeted to. After a scoped authorization, this list will be reduced to the
// spcific kube server used to authorize the request.
kubeServers []types.KubeServer
// isLocalKubernetesCluster is true if the target cluster is served by this teleport service.
// It is false if the target cluster is served by another teleport service or a different
// Teleport cluster.
isLocalKubernetesCluster bool
// LockingMode determines the kubernetes' behavior when locks are stale
LockingMode constants.LockingMode
}
func (c authContext) String() string {
return fmt.Sprintf("user: %v, users: %v, groups: %v, teleport cluster: %v, kube cluster: %v", c.User.GetName(), c.kubeUsers, c.kubeGroups, c.teleportCluster.name, c.kubeClusterName)
}
func (c *authContext) key() string {
// it is important that the context key contains user, kubernetes groups and certificate expiry,
// so that new logins with different parameters will not reuse this context
return fmt.Sprintf("%v:%v:%v:%v:%v:%v:%v", c.teleportCluster.name, c.User.GetName(), c.kubeUsers, c.kubeGroups, c.kubeClusterName, c.certExpires.Unix(), c.Identity.GetIdentity().ActiveRequests)
}
func (c *authContext) eventClusterMeta(req *http.Request) apievents.KubernetesClusterMetadata {
var kubeUsers, kubeGroups []string
if impersonateUser, impersonateGroups, err := computeAndValidateImpersonatedPrincipals(c.kubeUsers, c.kubeGroups, c.User.GetName(), req.Header); err == nil {
kubeUsers = []string{impersonateUser}
kubeGroups = impersonateGroups
} else {
kubeUsers = slices.Collect(maps.Keys(c.kubeUsers))
kubeGroups = slices.Collect(maps.Keys(c.kubeGroups))
}
return apievents.KubernetesClusterMetadata{
KubernetesCluster: c.kubeClusterName,
KubernetesUsers: kubeUsers,
KubernetesGroups: kubeGroups,
KubernetesLabels: c.kubeClusterLabels,
}
}
func (c *authContext) eventUserMeta() apievents.UserMetadata {
name := c.User.GetName()
meta := c.Identity.GetIdentity().GetUserMetadata()
meta.User = name
meta.Login = name
return meta
}
func (c *authContext) eventUserMetaWithLogin(login string) apievents.UserMetadata {
meta := c.eventUserMeta()
meta.Login = login
return meta
}
// getAccessChecker returns the [*services.ScopedAccessChecker] that granted access to the kube cluster
// refrenced by the authContext. Once it's set it should not be changed and all subsequent calls
// to an access checker should use the cached checker.
func (c *authContext) getAccessChecker() (*services.ScopedAccessChecker, error) {
if c.checker != nil {
return c.checker, nil
}
return nil, trace.AccessDenied("no access checker found for kube forwarder auth context")
}
// teleportClusterClient is a client for either a k8s endpoint in local cluster or a
// proxy endpoint in a remote cluster.
type teleportClusterClient struct {
remoteAddr utils.NetAddr
name string
isRemote bool
}
// handlerWithAuthFunc is http handler with passed auth context
type handlerWithAuthFunc func(ctx *authContext, w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error)
// handlerWithAuthFuncStd is http handler with passed auth context
type handlerWithAuthFuncStd func(ctx *authContext, w http.ResponseWriter, r *http.Request) (any, error)
// accessDeniedMsg is a message returned to the client when access is denied.
const accessDeniedMsg = "[00] access denied"
// authenticate function authenticates request
func (f *Forwarder) authenticate(req *http.Request) (*authContext, error) {
// If the cluster is not licensed for Kubernetes, return an error to the client.
if !f.cfg.ClusterFeatures.GetEntitlement(entitlements.K8s).Enabled {
// If the cluster is not licensed for Kubernetes, return an error to the client.
return nil, trace.AccessDenied("Teleport cluster is not licensed for Kubernetes")
}
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/authenticate",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
var isRemoteUser bool
userTypeI, err := authz.UserFromContext(ctx)
if err != nil {
f.log.WarnContext(ctx, "error getting user from context", "error", err)
return nil, trace.AccessDenied("%s", accessDeniedMsg)
}
switch userTypeI.(type) {
case authz.LocalUser:
case authz.RemoteUser:
isRemoteUser = true
case authz.BuiltinRole:
f.log.WarnContext(ctx, "Denying proxy access to unauthenticated user - this can sometimes be caused by inadvertently using an HTTP load balancer instead of a TCP load balancer on the Kubernetes port",
"user_type", logutils.TypeAttr(userTypeI),
)
return nil, trace.AccessDenied("%s", accessDeniedMsg)
default:
f.log.WarnContext(ctx, "Denying proxy access to unsupported user type", "user_type", logutils.TypeAttr(userTypeI))
return nil, trace.AccessDenied("%s", accessDeniedMsg)
}
scopedCtx, err := f.cfg.ScopedAuthz.AuthorizeScoped(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
authContext, err := f.setupContext(ctx, scopedCtx, req, isRemoteUser)
if err != nil {
f.log.WarnContext(ctx, "Unable to setup context", "error", err)
if trace.IsAccessDenied(err) {
if errors.Is(err, errAmbiguousCluster) {
return nil, trace.Wrap(err)
}
return nil, trace.AccessDenied("%s", accessDeniedMsg)
}
return nil, trace.Wrap(err)
}
return authContext, nil
}
func (f *Forwarder) withAuthStd(handler handlerWithAuthFuncStd) http.HandlerFunc {
return httplib.MakeStdHandlerWithErrorWriter(func(w http.ResponseWriter, req *http.Request) (any, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/withAuthStd",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
req = req.WithContext(ctx)
defer span.End()
authContext, err := f.authenticate(req)
if err != nil {
return nil, trace.Wrap(err)
}
if err := f.authorize(ctx, authContext); err != nil {
return nil, trace.Wrap(err)
}
return handler(authContext, w, req)
}, f.formatStatusResponseError)
}
// acquireConnectionLockWithIdentity acquires a connection lock under a given identity.
func (f *Forwarder) acquireConnectionLockWithIdentity(ctx context.Context, identity *authContext) error {
unscopedContext, isUnscoped := identity.UnscopedContext()
if !isUnscoped {
// TODO(espadolini) TODO(eriktate): scoped identities don't currently
// support max_kubernetes_connections, this should be updated when they
// do
return nil
}
maxConnections := unscopedContext.Checker.MaxKubernetesConnections()
if maxConnections == 0 {
return nil
}
user := unscopedContext.Identity.GetIdentity().Username
if err := f.acquireConnectionLock(ctx, user, maxConnections); err != nil {
return trace.Wrap(err)
}
return nil
}
// authOption is a functional option for authOptions.
type authOption func(*authOptions)
// authOptions is a set of options for withAuth handler.
type authOptions struct {
// errFormater is a function that formats the error response.
errFormater func(http.ResponseWriter, error)
}
// withCustomErrFormatter allows to override the default error formatter.
func withCustomErrFormatter(f func(http.ResponseWriter, error)) authOption {
return func(o *authOptions) {
o.errFormater = f
}
}
func (f *Forwarder) withAuth(handler handlerWithAuthFunc, opts ...authOption) httprouter.Handle {
authOpts := authOptions{
errFormater: f.formatStatusResponseError,
}
for _, opt := range opts {
opt(&authOpts)
}
return httplib.MakeHandlerWithErrorWriter(func(w http.ResponseWriter, req *http.Request, p httprouter.Params) (any, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/withAuth",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
req = req.WithContext(ctx)
defer span.End()
authContext, err := f.authenticate(req)
if err != nil {
return nil, trace.Wrap(err)
}
if err := f.authorize(ctx, authContext); err != nil {
return nil, trace.Wrap(err)
}
err = f.acquireConnectionLockWithIdentity(ctx, authContext)
if err != nil {
return nil, trace.Wrap(err)
}
return handler(authContext, w, req, p)
}, authOpts.errFormater)
}
// withAuthPassthrough authenticates the request and fetches information but doesn't deny if the user
// doesn't have RBAC access to the Kubernetes cluster.
func (f *Forwarder) withAuthPassthrough(handler handlerWithAuthFunc) httprouter.Handle {
return httplib.MakeHandlerWithErrorWriter(func(w http.ResponseWriter, req *http.Request, p httprouter.Params) (any, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/withAuthPassthrough",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
req = req.WithContext(ctx)
defer span.End()
authContext, err := f.authenticate(req)
if err != nil {
return nil, trace.Wrap(err)
}
err = f.acquireConnectionLockWithIdentity(req.Context(), authContext)
if err != nil {
return nil, trace.Wrap(err)
}
return handler(authContext, w, req, p)
}, f.formatStatusResponseError)
}
func (f *Forwarder) formatForwardResponseError(rw http.ResponseWriter, r *http.Request, respErr error) {
f.formatStatusResponseError(rw, respErr)
}
// writeResponseErrorToBody writes the error response to the body without any formatting.
// It is used for the /version endpoint since Kubernetes doesn't expect a JSON response
// for that endpoint.
func (f *Forwarder) writeResponseErrorToBody(rw http.ResponseWriter, respErr error) {
http.Error(rw, respErr.Error(), http.StatusInternalServerError)
}
// formatForwardResponseError handles errors returned from requests to the Kubernetes API.
// Any errors produced as a result of a GOAWAY request are forwarded to users as [http.StatusTooManyRequests]
// with a Retry-After header set to inform clients that they should retry the request. All
// other errors are formatted as a [metav1.Status] and written to the [http.ResponseWriter].
func (f *Forwarder) formatStatusResponseError(rw http.ResponseWriter, respErr error) {
// This detects failed requests that were terminated by the server due to GOAWAY. There
// is no direct way to detect these errors. No exported constants or error types exist from the
// standard library, so we have to match on the error message. The two error strings come from:
// - golang.org/x/net/http2 when its internal retry path cannot replay the body:
// https://github.com/golang/net/blob/5ac9daca088ab4f378d7df849f6c7d28bea86071/http2/transport.go#L694
// - net/http (errCannotRewind) when, after the http2 conn pool is drained, the http1 retry
// path tries to rewind the body and fails because Request.GetBody is unset:
// https://github.com/golang/go/blob/go1.26.2/src/net/http/transport.go#L759
// When a failed request is found, we return a response that indicates to clients that they
// should retry the request themselves.
errString := respErr.Error()
isHTTP2RetryErr := strings.Contains(errString, `http2: Transport: cannot retry err`) &&
strings.HasSuffix(errString, `after Request.Body was written; define Request.GetBody to avoid this error`)
isHTTP1RewindErr := strings.Contains(errString, `net/http: cannot rewind body after connection loss`)
if isHTTP2RetryErr || isHTTP1RewindErr {
data, err := runtime.Encode(globalKubeCodecs.LegacyCodec(), &kubeerrors.NewTooManyRequests("Connection closed by upstream Kubernetes server", 1).ErrStatus)
if err != nil {
f.log.WarnContext(f.ctx, "Failed encoding error into kube Status object", "error", err)
trace.WriteError(rw, respErr)
return
}
rw.Header().Set("Retry-After", "1")
rw.Header().Set(responsewriters.ContentTypeHeader, "application/json")
rw.WriteHeader(http.StatusTooManyRequests)
if _, err := rw.Write(data); err != nil && !utils.IsOKNetworkError(err) {
f.log.WarnContext(f.ctx, "Failed writing kube error response body", "error", err)
}
return
}
code, reason := kubeStatusCodeAndReason(respErr)
status := &metav1.Status{
Status: metav1.StatusFailure,
// Don't trace.Unwrap the error, in case it was wrapped with a
// user-friendly message. The underlying root error is likely too
// low-level to be useful.
Message: respErr.Error(),
Code: int32(code),
Reason: reason,
}
data, err := runtime.Encode(globalKubeCodecs.LegacyCodec(), status)
if err != nil {
f.log.WarnContext(f.ctx, "Failed encoding error into kube Status object", "error", err)
trace.WriteError(rw, respErr)
return
}
rw.Header().Set(responsewriters.ContentTypeHeader, "application/json")
// Always write the correct error code in the response so kubectl can parse
// it correctly. If response code and status.Code drift, kubectl prints
// `Error from server (InternalError): an error on the server ("unknown")
// has prevented the request from succeeding`` instead of the correct reason.
rw.WriteHeader(code)
if _, err := rw.Write(data); err != nil && !utils.IsOKNetworkError(err) {
f.log.WarnContext(f.ctx, "Failed writing kube error response body", "error", err)
}
}
// kubeStatusCodeAndReason returns HTTP status code and Kubernetes status reason to use when surfacing error to user.
// Without this, trace.ErrorToCode falls back to 500 and rewrites the original 403 into an InternalError.
func kubeStatusCodeAndReason(respErr error) (int, metav1.StatusReason) {
var statusErr *kubeerrors.StatusError
if errors.As(respErr, &statusErr) && statusErr.ErrStatus.Code != 0 {
return int(statusErr.ErrStatus.Code), statusErr.ErrStatus.Reason
}
code := trace.ErrorToCode(respErr)
reason := errorToKubeStatusReason(respErr, code)
return code, reason
}
var errAmbiguousCluster = &trace.AccessDeniedError{Message: "could not disambiguate between two or more scoped kube clusters with the same name, please login with credentials for a narrower scope"}
func (f *Forwarder) setupContext(
ctx context.Context,
scopedCtx *authz.ScopedContext,
req *http.Request,
isRemoteUser bool,
) (*authContext, error) {
ctx, span := f.cfg.tracer.Start(
ctx,
"kube.Forwarder/setupContext",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
identity := scopedCtx.Identity.GetIdentity()
teleportClusterName := identity.RouteToCluster
if teleportClusterName == "" {
teleportClusterName = f.cfg.ClusterName
}
isRemoteCluster := f.cfg.ClusterName != teleportClusterName
if isRemoteCluster && isRemoteUser {
return nil, trace.AccessDenied("access denied: remote user can not access remote cluster")
}
var (
kubeServers []types.KubeServer
kubeResource metaResource
err error
)
kubeCluster := identity.KubernetesCluster
unscopedCtx, isUnscoped := scopedCtx.UnscopedContext()
// aliasing to isScoped because it's a bit easier to reason about
isScoped := !isUnscoped
// Only check k8s principals for local clusters.
//
// For remote clusters, everything will be remapped to new roles on the
// leaf and checked there.
if !isRemoteCluster {
kubeServers, err = f.getKubernetesServersForKubeCluster(ctx, kubeCluster)
if err != nil || len(kubeServers) == 0 {
return nil, trace.NotFound("Kubernetes cluster %q not found", kubeCluster)
}
if !isScoped {
// If the calling identity is not scoped but there are multiple kube servers present in different scopes
// (e.g. when running in Proxy mode) we won't be able to disambiguate between them. In that case, we
// should return an error suggesting that the user logs in with scoped credentials
if err := checkAmbiguousClusters(kubeServers); err != nil {
return nil, trace.Wrap(err)
}
}
}
isLocalKubernetesCluster := f.isLocalKubeCluster(isRemoteCluster, kubeCluster)
if isLocalKubernetesCluster {
kubeResource, err = f.parseResourceFromRequest(req, kubeCluster)
if err != nil {
return nil, trace.Wrap(err)
}
} else {
kubeResource.verb = kubeResource.requestedResource.getVerb(req)
}
netConfig, err := f.cfg.CachingAuthClient.GetClusterNetworkingConfig(f.ctx)
if err != nil {
return nil, trace.Wrap(err)
}
recordingConfig, err := f.cfg.CachingAuthClient.GetSessionRecordingConfig(f.ctx)
if err != nil {
return nil, trace.Wrap(err)
}
authPref, err := f.cfg.CachingAuthClient.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// These are the defaults for scoped identities. They will be adjusted
// in the authorize function once we know which role is being used to
// grant access.
sessionTTL := time.Hour
clientIdleTimeout := netConfig.GetClientIdleTimeout()
if !isScoped {
sessionTTL = unscopedCtx.Checker.AdjustSessionTTL(sessionTTL)
clientIdleTimeout = unscopedCtx.Checker.AdjustClientIdleTimeout(clientIdleTimeout)
}
return &authContext{
ScopedContext: scopedCtx,
clientIdleTimeout: clientIdleTimeout,
clientIdleTimeoutMessage: netConfig.GetClientIdleTimeoutMessage(),
sessionTTL: sessionTTL,
recordingConfig: recordingConfig,
kubeClusterName: kubeCluster,
certExpires: identity.Expires,
disconnectExpiredCert: scopedCtx.GetDisconnectCertExpiry(authPref),
teleportCluster: teleportClusterClient{
name: teleportClusterName,
remoteAddr: utils.NetAddr{AddrNetwork: "tcp", Addr: req.RemoteAddr},
isRemote: isRemoteCluster,
},
kubeServers: kubeServers,
metaResource: kubeResource,
isLocalKubernetesCluster: isLocalKubernetesCluster,
}, nil
}
// checkAmbiguousClusters accepts a list of kube servers that have already been filtered for a given cluster name
// and determines if there is any ambiguity about which scope we should route to.
func checkAmbiguousClusters(kubeServers []types.KubeServer) error {
var foundScope string
for i, ks := range kubeServers {
switch {
case i == 0:
foundScope = ks.GetScope()
case ks.GetScope() != foundScope:
return trace.Wrap(errAmbiguousCluster)
}
}
return nil
}
func (f *Forwarder) parseResourceFromRequest(req *http.Request, kubeClusterName string) (metaResource, error) {
switch f.cfg.KubeServiceType {
case LegacyProxyService:
if details, err := f.findKubeDetailsByClusterName(kubeClusterName); err == nil {
out, err := getResourceFromRequest(req, details)
return out, trace.Wrap(err)
}
// When the cluster is not being served by the local service, the LegacyProxy
// is working as a normal proxy and will forward the request to the remote
// service. When this happens, proxy won't enforce any Kubernetes RBAC rules
// and will forward the request as is to the remote service. The remote
// service will enforce RBAC rules and will return an error if the user is
// not authorized.
fallthrough
case ProxyService:
// When the service is acting as a proxy (ProxyService or LegacyProxyService
// if the local cluster wasn't found), the proxy will forward the request
// to the remote service without enforcing any RBAC rules - we send the
// details = nil to indicate that we don't want to extract the kube resource
// from the request.
out, err := getResourceFromRequest(req, nil /*details*/)
return out, trace.Wrap(err)
case KubeService:
details, err := f.findKubeDetailsByClusterName(kubeClusterName)
if err != nil {
return metaResource{}, trace.Wrap(err)
}
out, err := getResourceFromRequest(req, details)
return out, trace.Wrap(err)
default:
return metaResource{}, trace.BadParameter("unsupported kube service type: %q", f.cfg.KubeServiceType)
}
}
// emitAuditEvent emits the audit event for a `kube.request` event if the session
// requires audit events.
func (f *Forwarder) emitAuditEvent(req *http.Request, sess *clusterSession, status int) {
_, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/emitAuditEvent",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
// If the session is not local, don't emit the event.
if !sess.isLocalKubernetesCluster {
return
}
r := sess.metaResource.requestedResource
if r.skipEvent {
return
}
// Emit audit event.
event := &apievents.KubeRequest{
Metadata: apievents.Metadata{
Type: events.KubeRequestEvent,
Code: events.KubeRequestCode,
},
UserMetadata: sess.eventUserMeta(),
ConnectionMetadata: apievents.ConnectionMetadata{
RemoteAddr: req.RemoteAddr,
LocalAddr: sess.kubeAddress,
Protocol: events.EventProtocolKube,
},
ServerMetadata: sess.getServerMetadata(),
RequestPath: req.URL.Path,
Verb: req.Method,
ResponseCode: int32(status),
KubernetesClusterMetadata: sess.eventClusterMeta(req),
SessionMetadata: apievents.SessionMetadata{
WithMFA: sess.Identity.GetIdentity().MFAVerified,
},
}
r.populateEvent(event)
if err := f.cfg.Emitter.EmitAuditEvent(f.ctx, event); err != nil {
f.log.WarnContext(f.ctx, "Failed to emit event", "error", err)
}
}
// fillDefaultKubePrincipalDetails fills the default details in order to keep
// the correct behavior when forwarding the request to the Kubernetes API.
// By default, if no kubernetes_users are set (which will be a majority), a
// user will impersonate himself, which is the backwards-compatible behavior.
// We also append teleport.KubeSystemAuthenticated to kubernetes_groups, which is
// a builtin group that allows any user to access common API methods,
// e.g. discovery methods required for initial client usage, without it,
// restricted user's kubectl clients will not work.
func fillDefaultKubePrincipalDetails(kubeUsers []string, kubeGroups []string, username string) ([]string, []string) {
if len(kubeUsers) == 0 {
kubeUsers = append(kubeUsers, username)
}
if !slices.Contains(kubeGroups, teleport.KubeSystemAuthenticated) {
kubeGroups = append(kubeGroups, teleport.KubeSystemAuthenticated)
}
return kubeUsers, kubeGroups
}
// kubeAccessDetails holds the allowed kube groups/users names and the cluster labels for a local kube cluster.
type kubeAccessDetails struct {
// list of allowed kube users
kubeUsers []string
// list of allowed kube groups
kubeGroups []string
// kube cluster labels
clusterLabels map[string]string
// kubeCluster is the local kube cluster we're granting access to
kubeCluster types.KubeCluster
// checker is the scoped access checker that permitted access to this
// kubeAccessDetails
checker *services.ScopedAccessChecker
}
var errNoMatchingCluster = &trace.AccessDeniedError{Message: "no matching kube cluster found"}
var errImplicitDeny = &trace.AccessDeniedError{Message: "no roles grant access"}
// getKubeAccessDetails returns the allowed kube groups/users names and the cluster labels for a local kube cluster.
// It also returns the [*services.ScopedAccessChecker] that granted access to the returned details.
func (f *Forwarder) getKubeAccessDetails(
ctx context.Context,
actx *authContext,
matchers ...services.RoleMatcher,
) (kubeAccessDetails, error) {
// track explicit denies so we can decide how to return once all available kube servers have been visited
var explicitDenies []error
// track the reason for an implicit deny so we can handle each case appropriately
var implicitDenyOrNotFound error = errNoMatchingCluster
// We don't want to append directly to matchers and don't want to clone matchers on every iteration when we
// append the label matcher. Instead we allocate a copy of the matchers with an extra slot reserved for adding the
// label prior to calling GetGroupsAndUsers. If matchers is empty, this will result in a slice of length 1 which
// will have its 0th element set to the label matcher.
matchersWithLabelMatcher := make([]services.RoleMatcher, len(matchers)+1)
copy(matchersWithLabelMatcher, matchers)
var checker *services.ScopedAccessChecker
// Find requested kubernetes cluster name and get allowed kube users/groups names.
for _, s := range actx.kubeServers {
c := s.GetCluster()
if c.GetName() != actx.kubeClusterName {
continue
}
// Once we find at least one cluster that matches by name, we should no longer return errNoMatchingCluster.
if errors.Is(implicitDenyOrNotFound, errNoMatchingCluster) {
implicitDenyOrNotFound = errImplicitDeny
}
if err := actx.CheckerContext.Decision(ctx, c.GetScope(), func(check *services.ScopedAccessChecker) error {
if err := check.Kube().CheckAccessToCluster(c, actx.accessState, matchers...); err != nil {
return err
}
checker = check
return nil
}); err != nil {
if services.IsAccessExplicitlyDenied(err) {
explicitDenies = append(explicitDenies, err)
}
continue
}
// Creates a matcher that matches the cluster labels against `kubernetes_labels`
// defined for each user's role. If a role has no `kubernetes_labels` defined, this matcher will
// treat it as a wildcard deny. We don't want this when checking access to the cluster, but we do
// when fetching groups and users so we only include it after getCheckerForCluster.
labels := types.CombineLabels(nil, c.GetStaticLabels(), types.LabelsToV2(c.GetDynamicLabels()))
labelMatcher := services.NewKubernetesClusterLabelMatcher(
types.CombineLabels(nil, c.GetStaticLabels(), types.LabelsToV2(c.GetDynamicLabels())),
checker.AccessInfo().Username,
actx.CheckerContext.Traits(),
)
matchersWithLabelMatcher[len(matchersWithLabelMatcher)-1] = labelMatcher
// GetGroupsAndUsers returns the accumulated kubernetes groups and users that satisfy the provided matchers.
// For unscoped identities, this will return the groups and users attached to any role that matches the
// target kubernetes cluster. If a KubernetesResourceMatcher is included in the list of matchers, it will
// only return the groups and users attached to roles that also satisfy the desired kubernetes resource.
// For scoped identities, the groups and users will be sourced from the scoped role providing access to the
// resource without matching on any specific kubernetes resources. The users/groups will be forwarded to the
// kubernetes cluster as impersonation headers.
const overrideTTL = false
groups, users, err := checker.Kube().GetGroupsAndUsers(checker.AdjustSessionTTL(actx.sessionTTL), overrideTTL, matchersWithLabelMatcher...)
if err != nil {
if services.IsAccessExplicitlyDenied(err) {
explicitDenies = append(explicitDenies, err)
}
if trace.IsNotFound(err) {
implicitDenyOrNotFound = err
}
continue
}
return kubeAccessDetails{
kubeGroups: groups,
kubeUsers: users,
clusterLabels: labels,
kubeCluster: c,
checker: checker,
}, nil
}
// Any explicit denials mean the cluster was found, no checkers granted access,
// and at least one explicitly denied access.
if len(explicitDenies) > 0 {
return kubeAccessDetails{}, trace.NewAggregate(explicitDenies...)
}
// If there are no explicit denials, we're left with three remaining cases:
// 1. There was no cluster found matching actx.kubeClusterName (no matching cluster)
// 2. There were no access checkers granting access to the cluster (implicit deny)
// 3. An access checker that would have granted access returned no groups/users (not found)
// The correct error for each case should already be assigned to implicitDenyOrNotFound
return kubeAccessDetails{}, trace.Wrap(implicitDenyOrNotFound)
}
func (f *Forwarder) authorize(ctx context.Context, actx *authContext) error {
ctx, span := f.cfg.tracer.Start(
ctx,
"kube.Forwarder/authorize",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
unscopedCtx, isUnscoped := actx.UnscopedContext()
// aliasing to isScoped because it's a bit easier to reason about
isScoped := !isUnscoped
if actx.teleportCluster.isRemote {
if isScoped {
return trace.Wrap(services.ErrScopedIdentity, "authorizing with remote cluster")
}
// Authorization for a remote kube cluster will happen on the remote
// end (by their proxy), after that cluster has remapped used roles.
f.log.DebugContext(ctx, "Skipping authorization for a remote kubernetes cluster name",
"auth_context", logutils.StringerAttr(actx),
)
return nil
}
if actx.kubeClusterName == "" {
if isScoped {
return trace.Wrap(services.ErrScopedIdentity, "authorizing with remote cluster")
}
// This should only happen for remote clusters (filtered above), but
// check and report anyway.
f.log.DebugContext(ctx, "Skipping authorization due to unknown kubernetes cluster name",
"auth_context", logutils.StringerAttr(actx),
)
return nil
}
// A kind unknown to the cluster's discovery can't be matched against the
// role's kubernetes_resources rules, so reject it rather than forward it unenforced.
// It's reported as NotFound because the kind isn't served by the cluster,
// the same result a client would get talking to the API server.
if actx.metaResource.unsupportedResource {
return trace.NotFound(
"Kubernetes resource kind %q in API group %q is not known to cluster %q",
actx.metaResource.requestedResource.resourceKind,
actx.metaResource.requestedResource.apiGroup,
actx.kubeClusterName,
)
}
identity := actx.Identity.GetIdentity()
var err error
actx.accessState, err = actx.CheckerContext.AccessStateFromTLSIdentity(ctx, &identity, f.cfg.CachingAuthClient)
if err != nil {
return trace.Wrap(err)
}
notFoundMessage := fmt.Sprintf("kubernetes cluster %q not found", actx.kubeClusterName)
var roleMatchers services.RoleMatchers
if !actx.metaResource.isList {
if rbacResource := actx.metaResource.rbacResource(); rbacResource != nil {
notFoundMessage = f.kubeResourceDeniedAccessMsg(
actx.User.GetName(),
actx.metaResource.verb,
actx.metaResource.requestedResource,
)
// If the kubeResource is available, append an extra matcher that validates
// if the kubernetes resource is allowed by the user roles that satisfy the
// target cluster labels.
// Each role defines `kubernetes_resources` and when kubeResource is available,
// KubernetesResourceMatcher will match roles that statisfy the resources at the
// same time that ClusterLabelMatcher matches the role's "kubernetes_labels".
// The call to roles.CheckKubeGroupsAndUsers when both matchers are provided
// results in the intersection of roles that match the "kubernetes_labels" and
// roles that allow access to the desired "kubernetes_resource".
// If from the intersection results an empty set, the request is denied.
//
// requiredRBACResources returns one tuple per (resource, verb) the request needs.
// Most requests need exactly one, adding an ephemeral container needs both exec and the mutation verb.
isClusterWideResource := actx.metaResource.isClusterWideResource()
required := actx.metaResource.requiredRBACResources()
roleMatchers = make(services.RoleMatchers, 0, len(required))
for i := range required {
roleMatchers = append(roleMatchers,
services.NewKubernetesResourceMatcher(required[i], isClusterWideResource))
}
}
}
// check access to cluster, check signing TTL, and return a list of allowed logins for local cluster based on
// Kubernetes service labels.
kubeAccessDetails, err := f.getKubeAccessDetails(
ctx,
actx,
roleMatchers...,
)
if errors.Is(err, services.ErrTrustedDeviceRequired) {
return trace.Wrap(err)
}
// roles.CheckKubeGroupsAndUsers returns trace.NotFound if the user does
// does not have at least one configured kubernetes_users or kubernetes_groups.
if trace.IsNotFound(err) {
const errMsg = "Your user's Teleport role does not allow Kubernetes access." +
" Please ask cluster administrator to ensure your role has appropriate kubernetes_groups and kubernetes_users set."
return trace.NotFound("%s", errMsg)
}
if err != nil {
if actx.metaResource.resourceDefinition != nil {
return trace.AccessDenied("%s", notFoundMessage)
}
if !errors.Is(err, errNoMatchingCluster) {
return trace.AccessDenied("%s", accessDeniedMsg)
}
}
// fillDefaultKubePrincipalDetails fills the default details in order to keep
// the correct behavior when forwarding the request to the Kubernetes API.
kubeUsers, kubeGroups := fillDefaultKubePrincipalDetails(kubeAccessDetails.kubeUsers, kubeAccessDetails.kubeGroups, actx.User.GetName())
actx.kubeUsers = set.New(kubeUsers...)
actx.kubeGroups = set.New(kubeGroups...)
actx.kubeCluster = kubeAccessDetails.kubeCluster
actx.kubeClusterLabels = kubeAccessDetails.clusterLabels
// cache the access checker for later decisions
actx.checker = kubeAccessDetails.checker
if actx.checker == nil && !isScoped {
// Checker could be nil if the request was implicitly denied, but this is a valid state for
// unscoped identities. We fall back to the unscoped context's checker in this case.
actx.checker = services.NewScopedAccessCheckerFromUnscoped(unscopedCtx.Checker)
}
actx.sessionTTL = actx.checker.AdjustSessionTTL(actx.sessionTTL)
actx.clientIdleTimeout, err = actx.checker.Kube().AdjustClientIdleTimeout(actx.clientIdleTimeout)
if err != nil {
return trace.Wrap(err)
}
authPref, err := f.cfg.CachingAuthClient.GetAuthPreference(ctx)
if err != nil {
return trace.Wrap(err)
}
if actx.checker.Kube().AdjustDisconnectExpiredCert(authPref.GetDisconnectExpiredCert()) {
actx.disconnectExpiredCert = actx.ScopedContext.GetDisconnectCertExpiryTime()
} else {
actx.disconnectExpiredCert = time.Time{}
}
// For scoped roles, check for the locking mode here so that we can verify whether users can connect or not when not
// when the lock is stale.
if isScoped {
actx.LockingMode = actx.checker.Kube().LockingMode(authPref.GetLockingMode())
// TODO(williamo/scopes): Potentially delete this in favor of checking locks in lib/authz/scoped.go
if err := f.cfg.LockWatcher.CheckLockInForce(actx.LockingMode, actx.LockTargets()...); err != nil {
return trace.Wrap(err)
}
}
// If the user has active Access requests we need to validate that they allow
// the kubeResource.
// This is required because CheckAccess does not validate the subresource type.
// TODO(eriktate/scopes): scoped identities don't support access requests, so we skip
// these checks for now.
if !isScoped && !actx.metaResource.isList {
if rbacResource := actx.metaResource.rbacResource(); rbacResource != nil && len(unscopedCtx.Checker.GetAllowedResourceAccessIDs()) > 0 {
// GetKubeResources returns the allowed and denied Kubernetes resources
// for the user. Since we have active access requests, the allowed
// resources will be the list of pods that the user requested access to if he
// requested access to specific pods or the list of pods that his roles
// allow if the user requested access a kubernetes cluster. If the user
// did not request access to any Kubernetes resource type, the allowed
// list will be empty.
allowed, denied := unscopedCtx.Checker.GetKubeResources(actx.kubeCluster)
if result, err := matchKubernetesResource(*rbacResource, actx.metaResource.isClusterWideResource(), allowed, denied); err != nil || !result {
return trace.AccessDenied("%s", notFoundMessage)
}
}
}
// if we were granted access to any kube users or groups, we can safely proceed
if len(actx.kubeUsers) > 0 || len(actx.kubeGroups) > 0 {
return nil
}
// If we were not granted access to any kube users or groups, we need to check if we should bypass auth for a
// proxy-based kube cluster. We should only skip authZ for unscoped identities connecting to proxy-based kube
// clusters.
if actx.kubeClusterName == f.cfg.ClusterName {
if isScoped {
return trace.Wrap(services.ErrScopedIdentity, "scoped identities do not support kube clusters exposed directly from the Teleport Proxy Service, only clusters exposed by the Teleport Kubernetes Service are supported")
}
f.log.DebugContext(ctx, "Skipping authorization for proxy-based kubernetes cluster",
"auth_context", logutils.StringerAttr(actx),
)
return nil
}
return trace.AccessDenied("%s", notFoundMessage)
}
// matchKubernetesResource checks if the Kubernetes Resource does not match any
// entry from the deny list and matches at least one entry from the allowed list.
func matchKubernetesResource(resource types.KubernetesResource, isClusterWideResource bool, allowed, denied []types.KubernetesResource) (bool, error) {
// utils.KubeResourceMatchesRegex checks if the resource.Kind is strictly equal
// to each entry and validates if the Name and Namespace fields matches the
// regex allowed by each entry.
result, err := utils.KubeResourceMatchesRegex(resource, isClusterWideResource, denied, types.Deny)
if err != nil {
return false, trace.Wrap(err)
} else if result {
return false, nil
}
result, err = utils.KubeResourceMatchesRegex(resource, isClusterWideResource, allowed, types.Allow)
if err != nil {
return false, trace.Wrap(err)
}
return result, nil
}
// join joins an existing session over a websocket connection
func (f *Forwarder) join(ctx *authContext, w http.ResponseWriter, req *http.Request, p httprouter.Params) (resp any, err error) {
// Increment the request counter and the in-flight gauge.
joinSessionsRequestCounter.WithLabelValues(f.cfg.KubeServiceType).Inc()
joinSessionsInFlightGauge.WithLabelValues(f.cfg.KubeServiceType).Inc()
defer joinSessionsInFlightGauge.WithLabelValues(f.cfg.KubeServiceType).Dec()
f.log.DebugContext(req.Context(), "Joining session", "join_url", logutils.StringerAttr(req.URL))
sess, err := f.newClusterSession(req.Context(), *ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// sess.Close cancels the connection monitor context to release it sooner.
// When the server is under heavy load it can take a while to identify that
// the underlying connection is gone. This change prevents that and releases
// the resources as soon as we know the session is no longer active.
defer sess.close()
if err := f.setupForwardingHeaders(sess, req, false /* withImpersonationHeaders */); err != nil {
return nil, trace.Wrap(err)
}
if !sess.isLocalKubernetesCluster {
return f.remoteJoin(ctx, w, req, p, sess)
}
sessionIDString := p.ByName("session")
sessionID, err := uuid.Parse(sessionIDString)
if err != nil {
return nil, trace.Wrap(err)
}
session := f.getSession(sessionID)
if session == nil {
return nil, trace.NotFound("session %v not found", sessionID)
}
ws, err := f.upgrader.Upgrade(w, req, nil)
if err != nil {
return nil, trace.Wrap(err)
}
var stream *streamproto.SessionStream
// Close the stream when we exit to ensure no goroutines are leaked and
// to ensure the client gets a close message in case of an error.
defer func() {
if stream != nil {
stream.Close()
}
}()
if err := func() error {
stream, err = streamproto.NewSessionStream(ws, streamproto.ServerHandshake{MFARequired: session.PresenceEnabled})
if err != nil {
return trace.Wrap(err)
}
client := &websocketClientStreams{uuid.New(), stream}
party := newParty(*ctx, stream.Mode, client)
err = session.join(req.Context(), party, true /* emitSessionJoinEvent */)
if err != nil {
return trace.Wrap(err)
}
defer func() {
// Detach cancellation so leave's moderation rebalance still runs after the websocket closes.
leaveCtx := context.WithoutCancel(req.Context())
if _, err := session.leave(leaveCtx, party.ID); err != nil {
f.log.DebugContext(leaveCtx, "Participant was unable to leave session",
"participant_id", party.ID,
"session_id", session.id,
"error", err,
)
}
}()
select {
case <-stream.Done():
party.InformClose(trace.BadParameter("websocket connection closed"))
return nil
case err := <-party.closeC:
return trace.Wrap(err)
}
}(); err != nil {
writeErr := ws.WriteControl(gwebsocket.CloseMessage, gwebsocket.FormatCloseMessage(gwebsocket.CloseInternalServerErr, err.Error()), time.Now().Add(time.Second*10))
if writeErr != nil {
f.log.WarnContext(req.Context(), "Failed to send early-exit websocket close message", "error", writeErr)
}
}
return nil, nil
}
// getSession retrieves the session from in-memory database.
// If the session was not found, returns nil.
// This method locks f.mu.
func (f *Forwarder) getSession(id uuid.UUID) *session {
f.mu.Lock()
defer f.mu.Unlock()
return f.sessions[id]
}
// setSession sets the session into in-memory database.
// If the session was not found, returns nil.
// This method locks f.mu.
func (f *Forwarder) setSession(id uuid.UUID, sess *session) {
f.mu.Lock()
defer f.mu.Unlock()
f.sessions[id] = sess
}
// deleteSession removes a session.
// This method locks f.mu.
func (f *Forwarder) deleteSession(id uuid.UUID) {
f.mu.Lock()
defer f.mu.Unlock()
delete(f.sessions, id)
}
// remoteJoin forwards a join request to a remote cluster.
func (f *Forwarder) remoteJoin(ctx *authContext, w http.ResponseWriter, req *http.Request, p httprouter.Params, sess *clusterSession) (resp any, err error) {
hostID, err := f.getSessionHostID(req.Context(), ctx, p)
if err != nil {
return nil, trace.Wrap(err)
}
netDialer := sess.DialWithContext(withTargetHostID(hostID))
tlsConfig, impersonationHeaders, err := f.getTLSConfig(sess)
if err != nil {
return nil, trace.Wrap(err)
}
dialer := &gwebsocket.Dialer{
TLSClientConfig: tlsConfig,
NetDialContext: netDialer,
}
headers := http.Header{}
if impersonationHeaders {
if headers, err = internal.IdentityForwardingHeaders(req.Context(), headers); err != nil {
return nil, trace.Wrap(err)
}
}
url := "wss://" + req.URL.Host
if req.URL.Port() != "" {
url = url + ":" + req.URL.Port()
}
url = url + req.URL.Path
wsTarget, respTarget, err := dialer.DialContext(req.Context(), url, headers)
if err != nil {
if respTarget == nil {
return nil, trace.Wrap(err)
}
defer respTarget.Body.Close()
msg, err := io.ReadAll(respTarget.Body)
if err != nil {
return nil, trace.Wrap(err)
}
var obj map[string]any
if err := json.Unmarshal(msg, &obj); err != nil {
return nil, trace.Wrap(err)
}
return obj, trace.Wrap(err)
}
defer wsTarget.Close()
defer respTarget.Body.Close()
wsSource, err := f.upgrader.Upgrade(w, req, nil)
if err != nil {
return nil, trace.Wrap(err)
}
defer wsSource.Close()
wsProxy(req.Context(), f.log, wsSource, wsTarget)
return nil, nil
}
// getSessionHostID returns the host ID that controls the session being joined.
// If the session is remote, returns an empty string, otherwise returns the host ID
// from the session tracker.
func (f *Forwarder) getSessionHostID(ctx context.Context, authCtx *authContext, p httprouter.Params) (string, error) {
if authCtx.teleportCluster.isRemote {
return "", nil
}
session := p.ByName("session")
if session == "" {
return "", trace.BadParameter("missing session ID")
}
sess, err := f.cfg.AuthClient.GetSessionTracker(ctx, session)
if err != nil {
return "", trace.Wrap(err)
}
return sess.GetHostID(), nil
}
// wsProxy proxies a websocket connection between two clusters transparently to allow for
// remote joins.
func wsProxy(ctx context.Context, log *slog.Logger, wsSource *gwebsocket.Conn, wsTarget *gwebsocket.Conn) {
errS := make(chan error, 1)
errT := make(chan error, 1)
wg := &sync.WaitGroup{}
forwardConn := func(dst, src *gwebsocket.Conn, errc chan<- error) {
defer dst.Close()
defer src.Close()
for {
msgType, msg, err := src.ReadMessage()
if err != nil {
m := gwebsocket.FormatCloseMessage(gwebsocket.CloseNormalClosure, err.Error())
var e *gwebsocket.CloseError
if errors.As(err, &e) {
if e.Code != gwebsocket.CloseNoStatusReceived {
m = gwebsocket.FormatCloseMessage(e.Code, e.Text)
}
}
errc <- err
dst.WriteMessage(gwebsocket.CloseMessage, m)
break
}
err = dst.WriteMessage(msgType, msg)
if err != nil {
errc <- err
break
}
}
}
wg.Add(2)
go func() {
defer wg.Done()
forwardConn(wsSource, wsTarget, errS)
}()
go func() {
defer wg.Done()
forwardConn(wsTarget, wsSource, errT)
}()
var err error
var from, to string
select {
case err = <-errS:
from = "client"
to = "upstream"
case err = <-errT:
from = "upstream"
to = "client"
}
var websocketErr *gwebsocket.CloseError
if errors.As(err, &websocketErr) && websocketErr.Code == gwebsocket.CloseAbnormalClosure {
log.DebugContext(ctx, "websocket proxying failed", "src", from, "target", to, "error", err)
}
wg.Wait()
}
// acquireConnectionLock acquires a semaphore used to limit connections to the Kubernetes agent.
// The semaphore is releasted when the request is returned/connection is closed.
// Returns an error if a semaphore could not be acquired.
func (f *Forwarder) acquireConnectionLock(ctx context.Context, user string, maxConnections int64) error {
ctx, span := f.cfg.tracer.Start(
ctx,
"kube.Forwarder/acquireConnectionLock",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
if maxConnections == 0 {
return nil
}
_, err := services.AcquireSemaphoreLock(ctx, services.SemaphoreLockConfig{
Service: f.cfg.AuthClient,
Expiry: sessionMaxLifetime,
Params: types.AcquireSemaphoreRequest{
SemaphoreKind: types.SemaphoreKindKubernetesConnection,
SemaphoreName: user,
MaxLeases: maxConnections,
Holder: user,
},
})
if err != nil {
if strings.Contains(err.Error(), teleport.MaxLeases) {
err = trace.AccessDenied("too many concurrent kubernetes connections for user %q (max=%d)",
user,
maxConnections,
)
}
return trace.Wrap(err)
}
return nil
}
// execNonInteractive handles all exec sessions without a TTY.
func (f *Forwarder) execNonInteractive(ctx *authContext, req *http.Request, _ httprouter.Params, request remoteCommandRequest, proxy *remoteCommandProxy, sess *clusterSession) error {
canStart, err := f.canStartSessionAlone(ctx)
if err != nil {
return trace.Wrap(err)
}
if !canStart {
return trace.AccessDenied("insufficient permissions to launch non-interactive session")
}
eventPodMeta := request.eventPodMeta(request.context, sess.kubeAPICreds)
sessionStart := f.cfg.Clock.Now().UTC()
serverMetadata := sess.getServerMetadata()
sessionMetadata := ctx.Identity.GetIdentity().GetSessionMetadata(uuid.NewString())
connectionMetdata := apievents.ConnectionMetadata{
RemoteAddr: req.RemoteAddr,
LocalAddr: sess.kubeAddress,
Protocol: events.EventProtocolKube,
}
sessionStartEvent := &apievents.SessionStart{
Metadata: apievents.Metadata{
Type: events.SessionStartEvent,
Code: events.SessionStartCode,
ClusterName: f.cfg.ClusterName,
},
ServerMetadata: serverMetadata,
SessionMetadata: sessionMetadata,
UserMetadata: ctx.eventUserMeta(),
ConnectionMetadata: connectionMetdata,
KubernetesClusterMetadata: ctx.eventClusterMeta(req),
KubernetesPodMetadata: eventPodMeta,
InitialCommand: request.cmd,
SessionRecording: ctx.recordingConfig.GetMode(),
}
if err := f.cfg.Emitter.EmitAuditEvent(f.ctx, sessionStartEvent); err != nil {
f.log.WarnContext(f.ctx, "Failed to emit event", "error", err)
return trace.Wrap(err)
}
execEvent := &apievents.Exec{
Metadata: apievents.Metadata{
Type: events.ExecEvent,
ClusterName: f.cfg.ClusterName,
},
ServerMetadata: serverMetadata,
SessionMetadata: sessionMetadata,
UserMetadata: ctx.eventUserMeta(),
ConnectionMetadata: connectionMetdata,
CommandMetadata: apievents.CommandMetadata{
Command: strings.Join(request.cmd, " "),
},
KubernetesClusterMetadata: ctx.eventClusterMeta(req),
KubernetesPodMetadata: eventPodMeta,
}
defer func() {
if err := f.cfg.Emitter.EmitAuditEvent(f.ctx, execEvent); err != nil {
f.log.WarnContext(f.ctx, "Failed to emit exec event", "error", err)
}
sessionEndEvent := &apievents.SessionEnd{
Metadata: apievents.Metadata{
Type: events.SessionEndEvent,
Code: events.SessionEndCode,
ClusterName: f.cfg.ClusterName,
},
ServerMetadata: serverMetadata,
SessionMetadata: sessionMetadata,
UserMetadata: ctx.eventUserMeta(),
ConnectionMetadata: connectionMetdata,
Interactive: false,
StartTime: sessionStart,
EndTime: f.cfg.Clock.Now().UTC(),
KubernetesClusterMetadata: ctx.eventClusterMeta(req),
KubernetesPodMetadata: eventPodMeta,
InitialCommand: request.cmd,
SessionRecording: ctx.recordingConfig.GetMode(),
}
if err := f.cfg.Emitter.EmitAuditEvent(f.ctx, sessionEndEvent); err != nil {
f.log.WarnContext(f.ctx, "Failed to emit session end event", "error", err)
}
}()
executor, executorCleanup, err := f.getExecutor(sess, req)
if err != nil {
execEvent.Code = events.ExecFailureCode
execEvent.Error, execEvent.ExitCode = exitCode(err)
f.log.WarnContext(f.ctx, "Failed creating executor", "error", err)
return trace.Wrap(err)
}
defer executorCleanup()
streamOptions := proxy.options()
err = executor.StreamWithContext(req.Context(), streamOptions)
if err != nil {
execEvent.Code = events.ExecFailureCode
execEvent.Error, execEvent.ExitCode = exitCode(err)
f.log.WarnContext(f.ctx, "Executor failed while streaming", "error", err)
return trace.Wrap(err)
}
execEvent.Code = events.ExecCode
return nil
}
// canStartSessionAlone returns true if the user associated with authCtx
// is allowed to start a session without moderation.
func (f *Forwarder) canStartSessionAlone(authCtx *authContext) (bool, error) {
unscopedCtx, isUnscoped := authCtx.UnscopedContext()
// TODO(eriktate/scopes): scoped access does not currently support session moderation, so we always allow
// scoped identities to start a session alone. An unscoped identity connecting to a scoped kube agent should
// still expect moderated session policy to be enforced. We should revisit this once scoped moderated sessions are
// addressed more wholistically.
if !isUnscoped {
return true, nil
}
policySets := unscopedCtx.Checker.SessionPolicySets()
authorizer := moderation.NewSessionAccessEvaluator(policySets, types.KubernetesSessionKind, authCtx.User.GetName())
canStart, _, err := authorizer.FulfilledFor(nil)
if err != nil {
return false, trace.Wrap(err)
}
return canStart, nil
}
func exitCode(err error) (errMsg, code string) {
var (
kubeStatusErr = &kubeerrors.StatusError{}
kubeExecErr = kubeexec.CodeExitError{}
)
if errors.As(err, &kubeStatusErr) {
if kubeStatusErr.ErrStatus.Status == metav1.StatusSuccess {
return
}
errMsg = kubeStatusErr.ErrStatus.Message
if errMsg == "" {
errMsg = string(kubeStatusErr.ErrStatus.Reason)
}
code = strconv.Itoa(int(kubeStatusErr.ErrStatus.Code))
} else if errors.As(err, &kubeExecErr) {
if kubeExecErr.Err != nil {
errMsg = kubeExecErr.Err.Error()
}
code = strconv.Itoa(kubeExecErr.Code)
} else if err != nil {
errMsg = err.Error()
}
return
}
// exec forwards all exec requests to the target server, captures
// all output from the session
func (f *Forwarder) exec(authCtx *authContext, w http.ResponseWriter, req *http.Request, p httprouter.Params) (resp any, err error) {
// Increment the request counter and the in-flight gauge.
execSessionsRequestCounter.WithLabelValues(f.cfg.KubeServiceType).Inc()
execSessionsInFlightGauge.WithLabelValues(f.cfg.KubeServiceType).Inc()
defer execSessionsInFlightGauge.WithLabelValues(f.cfg.KubeServiceType).Dec()
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/exec",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("Exec"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
f.log.DebugContext(ctx, "Starting exec", "exec_url", logutils.StringerAttr(req.URL))
defer func() {
if err != nil {
f.log.DebugContext(ctx, "Exec request failed", "error", err)
}
}()
sess, err := f.newClusterSession(ctx, *authCtx)
if err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(ctx, "Failed to create cluster session", "error", err)
return nil, trace.Wrap(err)
}
// sess.Close cancels the connection monitor context to release it sooner.
// When the server is under heavy load it can take a while to identify that
// the underlying connection is gone. This change prevents that and releases
// the resources as soon as we know the session is no longer active.
defer sess.close()
sess.forwarder, err = f.makeSessionForwarder(sess)
if err != nil {
return nil, trace.Wrap(err)
}
q := req.URL.Query()
request := remoteCommandRequest{
podNamespace: p.ByName("podNamespace"),
podName: p.ByName("podName"),
containerName: q.Get("container"),
cmd: q["command"],
stdin: utils.AsBool(q.Get("stdin")),
stdout: utils.AsBool(q.Get("stdout")),
stderr: utils.AsBool(q.Get("stderr")),
tty: utils.AsBool(q.Get("tty")),
httpRequest: req,
httpResponseWriter: w,
context: ctx,
pingPeriod: f.cfg.ConnPingPeriod,
idleTimeout: sess.clientIdleTimeout,
onResize: func(remotecommand.TerminalSize) {},
}
if err := f.setupForwardingHeaders(sess, req, true /* withImpersonationHeaders */); err != nil {
return nil, trace.Wrap(err)
}
return upgradeRequestToRemoteCommandProxy(request,
func(proxy *remoteCommandProxy) error {
sess.sendErrStatus = proxy.writeStatus
if !sess.isLocalKubernetesCluster {
// We're forwarding this to another kubernetes_service instance or Teleport proxy, let it handle session recording.
return f.remoteExec(req, sess, proxy)
}
if !request.tty {
return f.execNonInteractive(authCtx, req, p, request, proxy, sess)
}
client := newKubeProxyClientStreams(proxy)
party := newParty(*authCtx, types.SessionPeerMode, client)
session, err := newSession(*authCtx, f, req, p, party, sess)
if err != nil {
return trace.Wrap(err)
}
f.setSession(session.id, session)
if err = session.join(ctx, party, true /* emitSessionJoinEvent */); err != nil {
return trace.Wrap(err)
}
err = <-party.closeC
// Detach cancellation so leave's moderation rebalance still runs after the client disconnects.
leaveCtx := context.WithoutCancel(ctx)
if _, errLeave := session.leave(leaveCtx, party.ID); errLeave != nil {
f.log.DebugContext(leaveCtx, "Participant was unable to leave session",
"participant_id", party.ID,
"session_id", session.id,
"error", errLeave,
)
}
return trace.Wrap(err)
},
)
}
// remoteExec forwards an exec request to a remote cluster.
func (f *Forwarder) remoteExec(req *http.Request, sess *clusterSession, proxy *remoteCommandProxy) error {
executor, executorCleanup, err := f.getExecutor(sess, req)
if err != nil {
f.log.WarnContext(req.Context(), "Failed creating executor", "error", err)
return trace.Wrap(err)
}
defer executorCleanup()
streamOptions := proxy.options()
err = executor.StreamWithContext(req.Context(), streamOptions)
if err != nil {
f.log.WarnContext(req.Context(), "Executor failed while streaming", "error", err)
}
return trace.Wrap(err)
}
// portForward starts port forwarding to the remote cluster
func (f *Forwarder) portForward(authCtx *authContext, w http.ResponseWriter, req *http.Request, p httprouter.Params) (any, error) {
// Increment the request counter and the in-flight gauge.
portforwardRequestCounter.WithLabelValues(f.cfg.KubeServiceType).Inc()
portforwardSessionsInFlightGauge.WithLabelValues(f.cfg.KubeServiceType).Inc()
defer portforwardSessionsInFlightGauge.WithLabelValues(f.cfg.KubeServiceType).Dec()
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/portForward",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("portForward"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
f.log.DebugContext(ctx, "Handling port forward request",
"request_url", logutils.StringerAttr(req.URL),
"request_headers", req.Header,
)
sess, err := f.newClusterSession(ctx, *authCtx)
if err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(ctx, "Failed to create cluster session", "error", err)
return nil, trace.Wrap(err)
}
// sess.Close cancels the connection monitor context to release it sooner.
// When the server is under heavy load it can take a while to identify that
// the underlying connection is gone. This change prevents that and releases
// the resources as soon as we know the session is no longer active.
defer sess.close()
sess.forwarder, err = f.makeSessionForwarder(sess)
if err != nil {
return nil, trace.Wrap(err)
}
if err := f.setupForwardingHeaders(sess, req, true /* withImpersonationHeaders */); err != nil {
f.log.DebugContext(ctx, "DENIED Port forward", "request_url", logutils.StringerAttr(req.URL))
return nil, trace.Wrap(err)
}
dialer, dialerCleanup, err := f.getPortForwardDialer(sess, req)
if err != nil {
return nil, trace.Wrap(err)
}
defer dialerCleanup()
auditSent := map[string]bool{} // Set of `addr`. Can be multiple ports on single call. Using bool to simplify the check.
var auditSentMu sync.Mutex
onPortForward := func(addr string, success bool) {
if !sess.isLocalKubernetesCluster {
return
}
auditSentMu.Lock()
isAuditSent := auditSent[addr]
if !isAuditSent {
auditSent[addr] = true
}
auditSentMu.Unlock()
if isAuditSent {
return
}
portForward := &apievents.PortForward{
Metadata: apievents.Metadata{
Type: events.PortForwardEvent,
Code: events.PortForwardCode,
},
UserMetadata: authCtx.eventUserMeta(),
ConnectionMetadata: apievents.ConnectionMetadata{
LocalAddr: sess.kubeAddress,
RemoteAddr: req.RemoteAddr,
Protocol: events.EventProtocolKube,
},
Addr: addr,
Status: apievents.Status{
Success: success,
},
KubernetesClusterMetadata: sess.eventClusterMeta(req),
KubernetesPodMetadata: apievents.KubernetesPodMetadata{
KubernetesPodNamespace: p.ByName("podNamespace"),
KubernetesPodName: p.ByName("podName"),
},
}
if !success {
portForward.Code = events.PortForwardFailureCode
}
if err := f.cfg.Emitter.EmitAuditEvent(f.ctx, portForward); err != nil {
f.log.WarnContext(ctx, "Failed to emit event", "error", err)
}
}
defer func() {
for addr := range auditSent {
portForward := &apievents.PortForward{
Metadata: apievents.Metadata{
Type: events.PortForwardEvent,
Code: events.PortForwardStopCode,
},
UserMetadata: authCtx.eventUserMeta(),
ConnectionMetadata: apievents.ConnectionMetadata{
LocalAddr: sess.kubeAddress,
RemoteAddr: req.RemoteAddr,
Protocol: events.EventProtocolKube,
},
Addr: addr,
KubernetesClusterMetadata: sess.eventClusterMeta(req),
KubernetesPodMetadata: apievents.KubernetesPodMetadata{
KubernetesPodNamespace: p.ByName("podNamespace"),
KubernetesPodName: p.ByName("podName"),
},
}
if err := f.cfg.Emitter.EmitAuditEvent(f.ctx, portForward); err != nil {
f.log.WarnContext(ctx, "Failed to emit event", "error", err)
}
}
}()
q := req.URL.Query()
request := portForwardRequest{
podNamespace: p.ByName("podNamespace"),
podName: p.ByName("podName"),
ports: q["ports"],
context: ctx,
httpRequest: req,
httpResponseWriter: w,
onPortForward: onPortForward,
targetDialer: dialer,
pingPeriod: f.cfg.ConnPingPeriod,
idleTimeout: sess.clientIdleTimeout,
}
f.log.DebugContext(ctx, "Starting port forwarding", "request", request)
err = runPortForwarding(request)
if err != nil {
return nil, trace.Wrap(err)
}
f.log.DebugContext(ctx, "Completed port forwarding", "request", request)
return nil, nil
}
// runPortForwarding checks if the request contains WebSocket upgrade headers and
// decides which protocol the client expects.
// Go client uses SPDY while other clients still require WebSockets.
// This function will run until the end of the execution of the request.
func runPortForwarding(req portForwardRequest) error {
switch {
case wsstream.IsWebSocketRequestWithTunnelingProtocol(req.httpRequest):
return trace.Wrap(runPortForwardingTunneledHTTPStreams(req))
case wsstream.IsWebSocketRequest(req.httpRequest):
return trace.Wrap(runPortForwardingWebSocket(req))
default:
return trace.Wrap(runPortForwardingHTTPStreams(req))
}
}
const (
// ImpersonateHeaderPrefix is K8s impersonation prefix for impersonation feature:
// https://kubernetes.io/docs/reference/access-authn-authz/authentication/#user-impersonation
ImpersonateHeaderPrefix = "Impersonate-"
// ImpersonateUserHeader is impersonation header for users
ImpersonateUserHeader = "Impersonate-User"
// ImpersonateGroupHeader is K8s impersonation header for user
ImpersonateGroupHeader = "Impersonate-Group"
// ImpersonationRequestDeniedMessage is access denied message for impersonation
ImpersonationRequestDeniedMessage = "impersonation request has been denied"
)
func (f *Forwarder) setupForwardingHeaders(sess *clusterSession, req *http.Request, withImpersonationHeaders bool) error {
if withImpersonationHeaders {
if err := setupImpersonationHeaders(sess, req.Header); err != nil {
return trace.Wrap(err)
}
}
// Setup scheme, override target URL to the destination address
req.URL.Scheme = "https"
req.RequestURI = req.URL.Path + "?" + req.URL.RawQuery
// We only have a direct host to provide when using local creds.
// Otherwise, use kube-teleport-proxy-alpn.teleport.cluster.local to pass TLS handshake and leverage TLS Routing.
req.URL.Host = fmt.Sprintf("%s%s", constants.KubeTeleportProxyALPNPrefix, constants.APIDomain)
if sess.kubeAPICreds != nil {
req.URL.Host = sess.kubeAPICreds.getTargetAddr()
}
// add origin headers so the service consuming the request on the other site
// is aware of where it came from
req.Header.Add("X-Forwarded-Proto", "https")
req.Header.Add("X-Forwarded-Host", req.Host)
req.Header.Add("X-Forwarded-Path", req.URL.Path)
req.Header.Add("X-Forwarded-For", req.RemoteAddr)
return nil
}
// setupImpersonationHeaders sets up Impersonate-User and Impersonate-Group headers
func setupImpersonationHeaders(sess *clusterSession, headers http.Header) error {
// If the request is remote or this instance is a proxy,
// do not set up impersonation headers.
if sess.teleportCluster.isRemote || sess.kubeAPICreds == nil {
return nil
}
impersonateUser, impersonateGroups, err := computeAndValidateImpersonatedPrincipals(sess.kubeUsers, sess.kubeGroups, sess.User.GetName(), headers)
if err != nil {
return trace.Wrap(err)
}
return replaceImpersonationHeaders(headers, impersonateUser, impersonateGroups)
}
func replaceImpersonationHeaders(headers http.Header, impersonateUser string, impersonateGroups []string) error {
headers.Set(ImpersonateUserHeader, impersonateUser)
// Make sure to overwrite the exiting headers, instead of appending to
// them.
headers.Del(ImpersonateGroupHeader)
for _, group := range impersonateGroups {
headers.Add(ImpersonateGroupHeader, group)
}
return nil
}
// copyImpersonationHeaders copies the impersonation headers from the source
// request to the destination request.
func copyImpersonationHeaders(dst, src http.Header) {
dst.Del(ImpersonateUserHeader)
dst.Del(ImpersonateGroupHeader)
for _, v := range src.Values(ImpersonateUserHeader) {
dst.Add(ImpersonateUserHeader, v)
}
for _, v := range src.Values(ImpersonateGroupHeader) {
dst.Add(ImpersonateGroupHeader, v)
}
}
// computeAndValidateImpersonatedPrincipals computes the intersection between the information
// received in the `Impersonate-User` and `Impersonate-Groups` headers and the
// allowed values. If the user didn't specify any user and groups to impersonate,
// Teleport will use every group the user is allowed to impersonate.
// This function also validates the impersonateUser and impersonateGroups against
// HTTP header field value requirements to prevent header injection attacks.
func computeAndValidateImpersonatedPrincipals(kubeUsers, kubeGroups map[string]struct{}, username string, headers http.Header) (string, []string, error) {
_, hasUserWildcard := kubeUsers[types.Wildcard]
var impersonateUser string
var impersonateGroups []string
for header, values := range headers {
if !strings.HasPrefix(header, "Impersonate-") {
continue
}
switch header {
case ImpersonateUserHeader:
if impersonateUser != "" {
return "", nil, trace.AccessDenied("%v, user already specified to %q", ImpersonationRequestDeniedMessage, impersonateUser)
}
if len(values) == 0 || len(values) > 1 {
return "", nil, trace.AccessDenied("%v, invalid user header %q", ImpersonationRequestDeniedMessage, values)
}
// when Kubernetes go-client sends impersonated groups it also sends the impersonated user.
// The issue arrises when the impersonated user was not defined and the user want to just impersonate
// a subset of his groups. In that case the request would fail because empty user is not on
// ctx.kubeUsers. If Teleport receives an empty impersonated user it will ignore it and later will fill it
// with the Teleport username.
if len(values[0]) == 0 {
continue
}
impersonateUser = values[0]
if _, ok := kubeUsers[impersonateUser]; !ok && !hasUserWildcard {
return "", nil, trace.AccessDenied("%v, user header %q is not allowed in roles", ImpersonationRequestDeniedMessage, impersonateUser)
}
case ImpersonateGroupHeader:
for _, group := range values {
if _, ok := kubeGroups[group]; !ok {
return "", nil, trace.AccessDenied("%v, group header %q value is not allowed in roles", ImpersonationRequestDeniedMessage, group)
}
impersonateGroups = append(impersonateGroups, group)
}
default:
return "", nil, trace.AccessDenied("%v, unsupported impersonation header %q", ImpersonationRequestDeniedMessage, header)
}
}
impersonateGroups = apiutils.Deduplicate(impersonateGroups)
// By default, if no kubernetes_users is set (which will be a majority),
// user will impersonate themselves, which is the backwards-compatible behavior.
//
// As long as at least one `kubernetes_users` is set, the forwarder will start
// limiting the list of users allowed by the client to impersonate.
//
// If the users' role set does not include actual user name, it will be rejected,
// otherwise there will be no way to exclude the user from the list).
//
// If the `kubernetes_users` role set includes only one user
// (quite frequently that's the real intent), teleport will default to it,
// otherwise it will refuse to select.
//
// This will enable the use case when `kubernetes_users` has just one field to
// link the user identity with the IAM role, for example `IAM#{{external.email}}`
//
if impersonateUser == "" {
if hasUserWildcard {
impersonateUser = username
} else {
switch len(kubeUsers) {
// this is currently not possible as kube users have at least one
// user (user name), but in case if someone breaks it, catch here
case 0:
return "", nil, trace.AccessDenied("assumed at least one user to be present")
// if there is deterministic choice, make it to improve user experience
case 1:
for user := range kubeUsers {
impersonateUser = user
break
}
default:
return "", nil, trace.AccessDenied(
"please select a user to impersonate, refusing to select a user due to several kubernetes_users set up for this user")
}
}
}
if len(impersonateGroups) == 0 {
for group := range kubeGroups {
impersonateGroups = append(impersonateGroups, group)
}
}
// Validate impersonateUser and impersonateGroups against HTTP header field value
// requirements to prevent header injection attacks.
// requirements in http://www.w3.org/Protocols/rfc2616/rfc2616-sec4.html#sec4.2
if !httpguts.ValidHeaderFieldValue(impersonateUser) {
return "", nil, trace.BadParameter("invalid impersonated user header value: %q", impersonateUser)
}
for _, group := range impersonateGroups {
if !httpguts.ValidHeaderFieldValue(group) {
return "", nil, trace.BadParameter("invalid impersonated group header value: %q", group)
}
}
return impersonateUser, impersonateGroups, nil
}
// catchAll forwards all HTTP requests to the target k8s API server
func (f *Forwarder) catchAll(authCtx *authContext, w http.ResponseWriter, req *http.Request) (any, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/catchAll",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("catchAll"),
semconv.RPCSystemKey.String("kube"),
),
)
req = req.WithContext(ctx)
defer span.End()
sess, err := f.newClusterSession(ctx, *authCtx)
if err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(ctx, "Failed to create cluster session", "error", err)
return nil, trace.Wrap(err)
}
// sess.Close cancels the connection monitor context to release it sooner.
// When the server is under heavy load it can take a while to identify that
// the underlying connection is gone. This change prevents that and releases
// the resources as soon as we know the session is no longer active.
defer sess.close()
sess.upgradeToHTTP2 = true
sess.forwarder, err = f.makeSessionForwarder(sess)
if err != nil {
return nil, trace.Wrap(err)
}
if err := f.setupForwardingHeaders(sess, req, true /* withImpersonationHeaders */); err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(ctx, "Failed to set up forwarding headers", "error", err)
return nil, trace.Wrap(err)
}
isLocalKubeCluster := sess.isLocalKubernetesCluster
isListRequest := authCtx.metaResource.verb == types.KubeVerbList
// Watch requests can be send to a single resource or to a collection of resources.
// isWatchingCollectionRequest is true when the request is a watch request and
// the resource is a collection of resources, e.g. /api/v1/pods?watch=true.
// authCtx.kubeResource is only set when the request targets a single resource.
isWatchingCollectionRequest := authCtx.metaResource.verb == types.KubeVerbWatch && authCtx.metaResource.isList
switch {
case isListRequest || isWatchingCollectionRequest:
return f.listResources(sess, w, req)
case authCtx.metaResource.verb == types.KubeVerbDeleteCollection && isLocalKubeCluster:
return f.deleteResourcesCollection(sess, w, req)
default:
rw := httplib.NewResponseStatusRecorder(w)
sess.forwarder.ServeHTTP(rw, req)
f.emitAuditEvent(req, sess, rw.Status())
return nil, nil
}
}
// getWebsocketRestConfig builds a [*rest.Config] configuration to be
// used when upgrading requests via websocket.
func (f *Forwarder) getWebsocketRestConfig(sess *clusterSession, req *http.Request) (_ *rest.Config, cleanup func(), _ error) {
tlsConfig, useImpersonation, err := f.getTLSConfig(sess)
if err != nil {
return nil, nil, trace.Wrap(err)
}
upgradeRoundTripper := NewWebsocketRoundTripperWithDialer(roundTripperConfig{
ctx: req.Context(),
log: f.log,
sess: sess,
dialWithContext: sess.DialWithContext(),
tlsConfig: tlsConfig,
originalHeaders: req.Header,
useIdentityForwarding: useImpersonation,
proxier: sess.getProxier(),
})
rt := http.RoundTripper(upgradeRoundTripper)
if sess.kubeAPICreds != nil {
var err error
rt, err = sess.kubeAPICreds.wrapTransport(rt)
if err != nil {
upgradeRoundTripper.Cleanup()
return nil, nil, trace.Wrap(err)
}
}
rt = tracehttp.NewTransport(rt)
cfg := &rest.Config{
// WrapTransport will replace default roundTripper created for the WebsocketExecutor
// and on successfully established connection we will set upgrader's websocket connection.
WrapTransport: func(baseRt http.RoundTripper) http.RoundTripper {
if wrt, ok := baseRt.(*kwebsocket.RoundTripper); ok {
upgradeRoundTripper.onConnected = func(wsConn *gwebsocket.Conn) {
wrt.Conn = wsConn
}
}
return rt
},
}
return cfg, upgradeRoundTripper.Cleanup, nil
}
func (f *Forwarder) getWebsocketExecutor(sess *clusterSession, req *http.Request) (_ remotecommand.Executor, cleanup func(), _ error) {
f.log.DebugContext(req.Context(), "Creating websocket remote executor for request",
"request_method", req.Method,
"request_uri", req.RequestURI,
)
cfg, wsCleanup, err := f.getWebsocketRestConfig(sess, req)
if err != nil {
return nil, nil, trace.Wrap(err, "unable to create websocket executor")
}
executor, err := remotecommand.NewWebSocketExecutor(cfg, req.Method, req.URL.String())
if err != nil {
wsCleanup()
return nil, nil, trace.Wrap(err, "unable to create websocket executor")
}
return executor, wsCleanup, nil
}
func isRelevantWebsocketError(err error) bool {
return err != nil && !strings.Contains(err.Error(), "next reader: EOF")
}
func (f *Forwarder) getExecutor(sess *clusterSession, req *http.Request) (_ remotecommand.Executor, cleanup func(), _ error) {
wsExec, wsCleanup, err := f.getWebsocketExecutor(sess, req)
if err != nil {
return nil, nil, trace.Wrap(err, "unable to create websocket executor")
}
spdyExec, spdyCleanup, err := f.getSPDYExecutor(sess, req)
if err != nil {
wsCleanup()
return nil, nil, trace.Wrap(err, "unable to create spdy executor")
}
executor, err := remotecommand.NewFallbackExecutor(
wsExec,
spdyExec,
func(err error) bool {
// If the error is a known upgrade failure, we can retry with the other protocol.
return httpstream.IsUpgradeFailure(err) ||
httpstream.IsHTTPSProxyError(err) ||
kubeerrors.IsForbidden(err) ||
isTeleportUpgradeFailure(err)
})
if err != nil {
wsCleanup()
spdyCleanup()
return nil, nil, trace.Wrap(err, "unable to create fallback executor")
}
return executor, func() { wsCleanup(); spdyCleanup() }, nil
}
func (f *Forwarder) getSPDYExecutor(sess *clusterSession, req *http.Request) (_ remotecommand.Executor, cleanup func(), _ error) {
f.log.DebugContext(req.Context(), "Creating SPDY remote executor for request",
"request_method", req.Method,
"request_uri", req.RequestURI,
)
tlsConfig, useImpersonation, err := f.getTLSConfig(sess)
if err != nil {
return nil, nil, trace.Wrap(err)
}
upgradeRoundTripper := NewSpdyRoundTripperWithDialer(roundTripperConfig{
ctx: req.Context(),
sess: sess,
dialWithContext: sess.DialWithContext(),
tlsConfig: tlsConfig,
pingPeriod: f.cfg.ConnPingPeriod,
originalHeaders: req.Header,
useIdentityForwarding: useImpersonation,
proxier: sess.getProxier(),
})
rt := http.RoundTripper(upgradeRoundTripper)
if sess.kubeAPICreds != nil {
var err error
rt, err = sess.kubeAPICreds.wrapTransport(rt)
if err != nil {
upgradeRoundTripper.Cleanup()
return nil, nil, trace.Wrap(err)
}
}
rt = tracehttp.NewTransport(rt)
executor, err := remotecommand.NewSPDYExecutorForTransports(
rt,
spdy.NewUpgraderForStreaming(upgradeRoundTripper),
req.Method,
req.URL,
)
if err != nil {
upgradeRoundTripper.Cleanup()
return nil, nil, trace.Wrap(err)
}
return executor, upgradeRoundTripper.Cleanup, nil
}
func (f *Forwarder) getPortForwardDialer(sess *clusterSession, req *http.Request) (_ httpstream.Dialer, cleanup func(), _ error) {
wsDialer, wsCleanup, err := f.getWebsocketDialer(sess, req)
if err != nil {
return nil, nil, trace.Wrap(err)
}
spdyDialer, spdyCleanup, err := f.getSPDYDialer(sess, req)
if err != nil {
wsCleanup()
return nil, nil, trace.Wrap(err)
}
return portforward.NewFallbackDialerForStreaming(wsDialer, spdyDialer, func(err error) bool {
// If the error is a known upgrade failure, we can retry with the other protocol.
return httpstream.IsUpgradeFailure(err) ||
httpstream.IsHTTPSProxyError(err) ||
kubeerrors.IsForbidden(err) ||
isTeleportUpgradeFailure(err)
}), func() { wsCleanup(); spdyCleanup() }, nil
}
// getSPDYDialer returns a dialer that can be used to upgrade the connection
// to SPDY protocol.
// SPDY is a deprecated protocol, but it is still used by kubectl to manage data streams.
// The dialer uses an HTTP1.1 connection to upgrade to SPDY.
func (f *Forwarder) getSPDYDialer(sess *clusterSession, req *http.Request) (_ httpstream.Dialer, cleanup func(), _ error) {
tlsConfig, useImpersonation, err := f.getTLSConfig(sess)
if err != nil {
return nil, nil, trace.Wrap(err)
}
req = createSPDYRequest(req, PortForwardProtocolV1Name)
upgradeRoundTripper := NewSpdyRoundTripperWithDialer(roundTripperConfig{
ctx: req.Context(),
sess: sess,
dialWithContext: sess.DialWithContext(),
tlsConfig: tlsConfig,
pingPeriod: f.cfg.ConnPingPeriod,
originalHeaders: req.Header,
useIdentityForwarding: useImpersonation,
proxier: sess.getProxier(),
})
rt := http.RoundTripper(upgradeRoundTripper)
if sess.kubeAPICreds != nil {
var err error
rt, err = sess.kubeAPICreds.wrapTransport(rt)
if err != nil {
upgradeRoundTripper.Cleanup()
return nil, nil, trace.Wrap(err)
}
}
client := &http.Client{
Transport: tracehttp.NewTransport(rt),
}
return spdy.NewDialerForStreaming(spdy.NewUpgraderForStreaming(upgradeRoundTripper), client, req.Method, req.URL),
upgradeRoundTripper.Cleanup,
nil
}
func (f *Forwarder) getWebsocketDialer(sess *clusterSession, req *http.Request) (_ httpstream.Dialer, cleanup func(), _ error) {
cfg, wsCleanup, err := f.getWebsocketRestConfig(sess, req)
if err != nil {
return nil, nil, trace.Wrap(err, "unable to retrieve *rest.Config for websocket")
}
dialer, err := portforward.NewSPDYOverWebsocketDialerForStreaming(req.URL, cfg)
return dialer, wsCleanup, trace.Wrap(err)
}
// createSPDYRequest modifies the passed request to remove
// WebSockets headers and add SPDY upgrade information, including
// spdy protocols acceptable to the client.
func createSPDYRequest(req *http.Request, spdyProtocols ...string) *http.Request {
clone := req.Clone(req.Context())
// Clean up the websocket headers from the http request.
clone.Header.Del(wsstream.WebSocketProtocolHeader)
clone.Header.Del("Sec-Websocket-Key")
clone.Header.Del("Sec-Websocket-Version")
clone.Header.Del(httpstream.HeaderUpgrade)
// Update the http request for an upstream SPDY upgrade.
clone.Method = "POST"
clone.Body = nil // Remove the request body which is unused.
clone.Header.Set(httpstream.HeaderUpgrade, httpstreamspdy.HeaderSpdy31)
clone.Header.Del(httpstream.HeaderProtocolVersion)
for i := range spdyProtocols {
clone.Header.Add(httpstream.HeaderProtocolVersion, spdyProtocols[i])
}
return clone
}
// clusterSession contains authenticated user session to the target cluster:
// x509 short lived credentials, forwarding proxies and other data
type clusterSession struct {
authContext
parent *Forwarder
// kubeAPICreds are the credentials used to authenticate to the Kubernetes API server.
// It is non-nil if the kubernetes cluster is served by this teleport service,
// nil otherwise.
kubeAPICreds kubeCreds
forwarder *reverseproxy.Forwarder
// targetAddr is the address of the target cluster.
targetAddr string
// kubeAddress is the address of this session's active connection (if there is one)
kubeAddress string
// upgradeToHTTP2 indicates whether the transport should be configured to use HTTP2.
// A HTTP2 configured transport does not work with connections that are going to be
// upgraded to SPDY, like in the cases of exec, port forward...
upgradeToHTTP2 bool
// requestContext is the context of the original request.
requestContext context.Context
// codecFactory is the codec factory used to create the serializer
// for unmarshalling the payload.
codecFactory *serializer.CodecFactory
// rbacSupportedResources is the list of resources that support RBAC for the
// current cluster.
rbacSupportedResources rbacSupportedResources
// sessionCtx is used with one or more connection contexts.
sessionCtx context.Context
// sessionCancel cancels the session context and related connection contexts.
sessionCancel context.CancelCauseFunc
// sendErrStatus is a function that sends an error status to the client.
sendErrStatus func(status *kubeerrors.StatusError) error
}
// close cancels the session context and related connection contexts.
func (s *clusterSession) close() {
s.sessionCancel(io.EOF)
}
func (s *clusterSession) monitorConn(conn net.Conn, err error, hostID string) (net.Conn, error) {
if err != nil {
return nil, trace.Wrap(err)
}
// Create a connection context from the session context.
// This separates session lifecycle from the connection attempt lifecycle.
// There may be multiple connection attempts within a session using FallbackExecutor/FallbackDialer.
// The approach avoids a potential race condition for s.sessionCancel.
connCtx, connCancel := context.WithCancelCause(s.sessionCtx)
tc, err := srv.NewTrackingReadConn(srv.TrackingReadConnConfig{
Conn: conn,
Clock: s.parent.cfg.Clock,
Context: connCtx,
Cancel: connCancel,
})
if err != nil {
connCancel(err)
return nil, trace.Wrap(err)
}
lockTargets := s.LockTargets()
// when the target is not a kubernetes_service instance, we don't need to lock it.
// the target could be a remote cluster or a local Kubernetes API server. In both cases,
// hostID is empty.
if hostID != "" {
lockTargets = append(lockTargets, types.LockTarget{
ServerID: hostID,
})
}
err = srv.StartMonitor(srv.MonitorConfig{
LockWatcher: s.parent.cfg.LockWatcher,
LockTargets: lockTargets,
DisconnectExpiredCert: s.disconnectExpiredCert,
ClientIdleTimeout: s.clientIdleTimeout,
IdleTimeoutMessage: s.clientIdleTimeoutMessage,
Clock: s.parent.cfg.Clock,
Tracker: tc,
Conn: tc,
Context: connCtx,
TeleportUser: s.User.GetName(),
UserOriginClusterName: s.Identity.GetIdentity().OriginClusterName,
ServerID: s.parent.cfg.HostID,
Logger: s.parent.log,
Emitter: s.parent.cfg.AuthClient,
EmitterContext: s.parent.ctx,
MessageWriter: formatForwardResponseError(s.sendErrStatus),
LockingMode: s.LockingMode,
})
if err != nil {
tc.CloseWithCause(err)
return nil, trace.Wrap(err)
}
return tc, nil
}
func (s *clusterSession) getServerMetadata() apievents.ServerMetadata {
return apievents.ServerMetadata{
ServerID: s.parent.cfg.HostID,
ServerNamespace: s.parent.cfg.Namespace,
ServerHostname: s.teleportCluster.name,
ServerAddr: s.kubeAddress,
ServerLabels: maps.Clone(s.kubeClusterLabels),
ServerVersion: teleport.Version,
}
}
func (s *clusterSession) Dial(network, addr string) (net.Conn, error) {
var hostID string
conn, err := s.dial(s.requestContext, network, addr, withHostIDCollection(&hostID))
return s.monitorConn(conn, err, hostID)
}
func (s *clusterSession) DialWithContext(opts ...contextDialerOption) func(ctx context.Context, network, addr string) (net.Conn, error) {
return func(ctx context.Context, network, addr string) (net.Conn, error) {
var hostID string
conn, err := s.dial(ctx, network, addr, append(opts, withHostIDCollection(&hostID))...)
return s.monitorConn(conn, err, hostID)
}
}
func (s *clusterSession) dial(ctx context.Context, network, addr string, opts ...contextDialerOption) (net.Conn, error) {
dialer := s.parent.getContextDialerFunc(s, opts...)
conn, err := dialer(ctx, network, addr)
return conn, trace.Wrap(err)
}
// getProxier returns the proxier function to use for this session.
// If the target cluster is not served by this teleport service, the proxier
// must be nil to avoid using it through the reverse tunnel.
// If the target cluster is served by this teleport service, the proxier
// must be set to the default proxy function.
func (s *clusterSession) getProxier() func(req *http.Request) (*url.URL, error) {
// When the target cluster is not served by this teleport service, the
// proxier must be nil to avoid using it through the reverse tunnel.
if s.kubeAPICreds == nil {
return nil
}
return utilnet.NewProxierWithNoProxyCIDR(http.ProxyFromEnvironment)
}
// getClusterScope returns the scope associated with the cluster that the clusterSession
// refers to.
func (s *clusterSession) getClusterScope() string {
if s.kubeCluster != nil {
return s.kubeCluster.GetScope()
}
return ""
}
// TODO(awly): unit test this
func (f *Forwarder) newClusterSession(ctx context.Context, authCtx authContext) (*clusterSession, error) {
ctx, span := f.cfg.tracer.Start(
ctx,
"kube.Forwarder/newClusterSession",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("GlobalRequest"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
if authCtx.teleportCluster.isRemote {
return f.newClusterSessionRemoteCluster(ctx, authCtx)
}
return f.newClusterSessionSameCluster(ctx, authCtx)
}
func (f *Forwarder) newClusterSessionRemoteCluster(ctx context.Context, authCtx authContext) (*clusterSession, error) {
f.log.DebugContext(ctx, "Forwarding kubernetes session to remote cluster", "auth_context", logutils.StringerAttr(authCtx))
sessionCtx, sessionCancel := context.WithCancelCause(ctx)
return &clusterSession{
parent: f,
authContext: authCtx,
// Proxy uses reverse tunnel dialer to connect to Kubernetes in a leaf cluster
// and the targetKubernetes cluster endpoint is determined from the identity
// encoded in the TLS certificate. We're setting the dial endpoint to a hardcoded
// `kube.teleport.cluster.local` value to indicate this is a Kubernetes proxy request
targetAddr: reversetunnelclient.LocalKubernetes,
requestContext: ctx,
sessionCtx: sessionCtx,
sessionCancel: sessionCancel,
}, nil
}
func (f *Forwarder) newClusterSessionSameCluster(ctx context.Context, authCtx authContext) (*clusterSession, error) {
// Try local creds first
sess, localErr := f.newClusterSessionLocal(ctx, authCtx)
switch {
case localErr == nil:
return sess, nil
case trace.IsConnectionProblem(localErr):
return nil, trace.Wrap(localErr)
}
kubeServers := authCtx.kubeServers
if len(kubeServers) == 0 && authCtx.kubeClusterName == authCtx.teleportCluster.name {
return nil, trace.Wrap(localErr)
}
if len(kubeServers) == 0 {
return nil, trace.NotFound("kubernetes cluster %q not found", authCtx.kubeClusterName)
}
return f.newClusterSessionDirect(ctx, authCtx)
}
func (f *Forwarder) newClusterSessionLocal(ctx context.Context, authCtx authContext) (*clusterSession, error) {
details, err := f.findKubeDetailsByClusterName(authCtx.kubeClusterName)
if err != nil {
return nil, trace.NotFound("kubernetes cluster %q not found", authCtx.kubeClusterName)
}
codecFactory, rbacSupportedResources, err := details.getClusterSupportedResources()
if err != nil {
return nil, trace.Wrap(err)
}
sessionCtx, sessionCancel := context.WithCancelCause(ctx)
f.log.DebugContext(ctx, "Handling kubernetes session using local credentials", "auth_context", logutils.StringerAttr(authCtx))
return &clusterSession{
parent: f,
authContext: authCtx,
kubeAPICreds: details.kubeCreds,
targetAddr: details.getTargetAddr(),
requestContext: ctx,
codecFactory: codecFactory,
rbacSupportedResources: rbacSupportedResources,
sessionCtx: sessionCtx,
sessionCancel: sessionCancel,
}, nil
}
func (f *Forwarder) newClusterSessionDirect(ctx context.Context, authCtx authContext) (*clusterSession, error) {
sessionCtx, sessionCancel := context.WithCancelCause(ctx)
return &clusterSession{
parent: f,
authContext: authCtx,
requestContext: ctx,
sessionCtx: sessionCtx,
sessionCancel: sessionCancel,
}, nil
}
// makeSessionForwader creates a new forward.Forwarder with a transport that
// is either configured:
// - for HTTP1 in case it's going to be used against streaming andoints like exec and port forward.
// - for HTTP2 in all other cases.
// The reason being is that streaming requests are going to be upgraded to SPDY, which is only
// supported coming from an HTTP1 request.
func (f *Forwarder) makeSessionForwarder(sess *clusterSession) (*reverseproxy.Forwarder, error) {
transport, err := f.transportForRequest(sess)
if err != nil {
return nil, trace.Wrap(err)
}
opts := []reverseproxy.Option{
reverseproxy.WithFlushInterval(100 * time.Millisecond),
reverseproxy.WithRoundTripper(transport),
reverseproxy.WithLogger(f.log),
reverseproxy.WithErrorHandler(f.formatForwardResponseError),
}
if sess.isLocalKubernetesCluster {
// If the target cluster is local, i.e. the cluster that is served by this
// teleport service, then we set up the forwarder to allow re-writing
// the response to the client to include user friendly error messages.
// This is done by adding a response modifier to the forwarder.
// Right now, the only error that is re-written is the 403 Forbidden error
// that is returned when the user tries to access a GKE Autopilot cluster
// with system:masters group impersonation.
//nolint:bodyclose // the caller closes the response body in httputils.ReverseProxy
opts = append(opts, reverseproxy.WithResponseModifier(f.rewriteResponseForbidden(sess)))
}
forwarder, err := reverseproxy.New(
opts...,
)
return forwarder, trace.Wrap(err)
}
// kubeClusters returns the list of available clusters
func (f *Forwarder) kubeClusters() types.KubeClusters {
f.rwMutexDetails.RLock()
defer f.rwMutexDetails.RUnlock()
res := make(types.KubeClusters, 0, len(f.clusterDetails))
for _, cred := range f.clusterDetails {
cluster := cred.kubeCluster.Copy()
res = append(res,
cluster,
)
}
return res
}
// findKubeDetailsByClusterName searches for the cluster details otherwise returns a trace.NotFound error.
func (f *Forwarder) findKubeDetailsByClusterName(name string) (*kubeDetails, error) {
f.rwMutexDetails.RLock()
defer f.rwMutexDetails.RUnlock()
if creds, ok := f.clusterDetails[name]; ok {
return creds, nil
}
return nil, trace.NotFound("cluster %s not found", name)
}
// upsertKubeDetails updates the details in f.ClusterDetails for key if they exist,
// otherwise inserts them.
func (f *Forwarder) upsertKubeDetails(key string, clusterDetails *kubeDetails) {
f.rwMutexDetails.Lock()
defer f.rwMutexDetails.Unlock()
if oldDetails, ok := f.clusterDetails[key]; ok {
oldDetails.Close()
}
// replace existing details in map
f.clusterDetails[key] = clusterDetails
}
// removeKubeDetails removes the kubeDetails from map.
func (f *Forwarder) removeKubeDetails(name string) {
f.rwMutexDetails.Lock()
defer f.rwMutexDetails.Unlock()
if oldDetails, ok := f.clusterDetails[name]; ok {
oldDetails.Close()
}
delete(f.clusterDetails, name)
}
// isLocalKubeCluster checks if the current service must hold the cluster and
// if it's of Type KubeService.
// KubeProxy services or remote clusters are automatically forwarded to
// the final destination.
func (f *Forwarder) isLocalKubeCluster(isRemoteTeleportCluster bool, kubeClusterName string) bool {
switch f.cfg.KubeServiceType {
case KubeService:
// Kubernetes service is always local.
return true
case LegacyProxyService:
// remote clusters are always forwarded to the final destination.
if isRemoteTeleportCluster {
return false
}
// Legacy proxy service is local only if the kube cluster name matches
// with clusters served by this agent.
_, err := f.findKubeDetailsByClusterName(kubeClusterName)
return err == nil
default:
return false
}
}
// kubeResourceDeniedAccessMsg creates a Kubernetes API like forbidden response.
// Logic from:
// https://github.com/kubernetes/kubernetes/blob/ea0764452222146c47ec826977f49d7001b0ea8c/staging/src/k8s.io/apiserver/pkg/endpoints/handlers/responsewriters/errors.go#L51
func (f *Forwarder) kubeResourceDeniedAccessMsg(user, verb string, resource apiResource) string {
kind := strings.Split(resource.resourceKind, "/")[0]
apiGroup := resource.apiGroup
teleportType := resource.resourceKind
switch {
case resource.namespace != "" && resource.resourceName != "":
// <resource> "<pod_name>" is forbidden: User "<user>" cannot create resource "<resource>" in API group "" in the namespace "<namespace>"
return fmt.Sprintf(
"%[1]s %[2]q is forbidden: User %[3]q cannot %[4]s resource %[1]q in API group %[5]q in the namespace %[6]q\n"+
"Ask your Teleport admin to ensure that your Teleport role includes access to the %[7]s in %[8]q field.\n"+
"Check by running: kubectl auth can-i %[4]s %[1]s/%[2]s --namespace %[6]s ",
kind, // 1
resource.resourceName, // 2
user, // 3
verb, // 4
apiGroup, // 5
resource.namespace, // 6
teleportType, // 7
kubernetesResourcesKey, // 8
)
case resource.namespace != "":
// <resource> is forbidden: User "<user>" cannot create resource "<resource>" in API group "" in the namespace "<namespace>"
return fmt.Sprintf(
"%[1]s is forbidden: User %[2]q cannot %[3]s resource %[1]q in API group %[4]q in the namespace %[5]q\n"+
"Ask your Teleport admin to ensure that your Teleport role includes access to the %[6]s in %[7]q field.\n"+
"Check by running: kubectl auth can-i %[3]s %[1]s --namespace %[5]s ",
kind, // 1
user, // 2
verb, // 3
apiGroup, // 4
resource.namespace, // 5
teleportType, // 6
kubernetesResourcesKey, // 7
)
case resource.resourceName == "":
return fmt.Sprintf(
"%[1]s is forbidden: User %[2]q cannot %[3]s resource %[1]q in API group %[4]q at the cluster scope\n"+
"Ask your Teleport admin to ensure that your Teleport role includes access to the %[5]s in %[6]q field.\n"+
"Check by running: kubectl auth can-i %[3]s %[1]s",
kind, // 1
user, // 2
verb, // 3
apiGroup, // 4
teleportType, // 5
kubernetesResourcesKey, // 6
)
default:
return fmt.Sprintf(
"%[1]s %[2]q is forbidden: User %[3]q cannot %[4]s resource %[1]q in API group %[5]q at the cluster scope\n"+
"Ask your Teleport admin to ensure that your Teleport role includes access to the %[6]s in %[7]q field.\n"+
"Check by running: kubectl auth can-i %[4]s %[1]s/%[2]s",
kind, // 1
resource.resourceName, // 2
user, // 3
verb, // 4
apiGroup, // 5
teleportType, // 6
kubernetesResourcesKey, // 7
)
}
}
// errorToKubeStatusReason returns an appropriate StatusReason based on the
// provided error type.
func errorToKubeStatusReason(err error, code int) metav1.StatusReason {
switch {
case trace.IsAggregate(err):
return metav1.StatusReasonTimeout
case trace.IsNotFound(err):
return metav1.StatusReasonNotFound
case trace.IsBadParameter(err) || trace.IsOAuth2(err):
return metav1.StatusReasonBadRequest
case trace.IsNotImplemented(err):
return metav1.StatusReasonMethodNotAllowed
case trace.IsCompareFailed(err):
return metav1.StatusReasonConflict
case trace.IsAccessDenied(err):
return metav1.StatusReasonForbidden
case trace.IsAlreadyExists(err):
return metav1.StatusReasonConflict
case trace.IsLimitExceeded(err):
return metav1.StatusReasonTooManyRequests
case trace.IsConnectionProblem(err):
return metav1.StatusReasonTimeout
case code == http.StatusInternalServerError:
return metav1.StatusReasonInternalError
default:
return metav1.StatusReasonUnknown
}
}
// formatForwardResponseError formats the error response from the connection
// monitor to a Kubernetes API error response.
type formatForwardResponseError func(status *kubeerrors.StatusError) error
func (f formatForwardResponseError) WriteString(s string) (int, error) {
if f == nil {
return len(s), nil
}
err := f(
&kubeerrors.StatusError{
ErrStatus: metav1.Status{
Status: metav1.StatusFailure,
Code: http.StatusInternalServerError,
Reason: metav1.StatusReasonInternalError,
Message: s,
},
},
)
if err != nil {
return 0, trace.Wrap(err)
}
return len(s), nil
}
// allHTTPMethods returns a list of all HTTP methods, useful for creating
// non-root catch-all handlers.
func allHTTPMethods() []string {
return []string{
http.MethodConnect,
http.MethodDelete,
http.MethodGet,
http.MethodHead,
http.MethodOptions,
http.MethodPatch,
http.MethodPost,
http.MethodPut,
http.MethodTrace,
}
}
/*
* 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 proxy
import (
"context"
"crypto/tls"
"log/slog"
"net/http"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"k8s.io/client-go/kubernetes"
"k8s.io/client-go/rest"
"k8s.io/client-go/transport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/services"
)
type kubeCreds interface {
getTLSConfig() *tls.Config
getTransportConfig() *transport.Config
getTargetAddr() string
getKubeRestConfig() *rest.Config
getKubeClient() kubernetes.Interface
getTransport() http.RoundTripper
wrapTransport(http.RoundTripper) (http.RoundTripper, error)
close() error
}
var (
_ kubeCreds = &staticKubeCreds{}
_ kubeCreds = &dynamicKubeCreds{}
)
// staticKubeCreds contain authentication-related fields from kubeconfig.
//
// TODO(awly): make this an interface, one implementation for local k8s cluster
// and another for a remote teleport cluster.
type staticKubeCreds struct {
// tlsConfig contains (m)TLS configuration.
tlsConfig *tls.Config
// transportConfig contains HTTPS-related configuration.
// Note: use wrapTransport method if working with http.RoundTrippers.
transportConfig *transport.Config
// targetAddr is a kubernetes API address.
targetAddr string
kubeClient kubernetes.Interface
// clientRestCfg is the Kubernetes Rest config for the cluster.
clientRestCfg *rest.Config
transport http.RoundTripper
}
func (s *staticKubeCreds) getTLSConfig() *tls.Config {
return s.tlsConfig.Clone()
}
func (s *staticKubeCreds) getTransport() http.RoundTripper {
return s.transport
}
func (s *staticKubeCreds) getTransportConfig() *transport.Config {
return s.transportConfig
}
func (s *staticKubeCreds) getTargetAddr() string {
return s.targetAddr
}
func (s *staticKubeCreds) getKubeClient() kubernetes.Interface {
return s.kubeClient
}
func (s *staticKubeCreds) getKubeRestConfig() *rest.Config {
return s.clientRestCfg
}
func (s *staticKubeCreds) wrapTransport(rt http.RoundTripper) (http.RoundTripper, error) {
if s == nil {
return rt, nil
}
wrapped, err := transport.HTTPWrappersForConfig(s.transportConfig, rt)
if err != nil {
return nil, trace.Wrap(err)
}
return enforceCloseIdleConnections(wrapped, rt), nil
}
// enforceCloseIdleConnections ensures that the returned [http.RoundTripper]
// has a CloseIdleConnections method. [transport.HTTPWrappersForConfig] returns
// a [http.RoundTripper] that does not implement it so any calls to [http.Client.CloseIdleConnections]
// will result in a noop instead of forwarding the request onto its wrapped [http.RoundTripper].
func enforceCloseIdleConnections(wrapper, wrapped http.RoundTripper) http.RoundTripper {
type closeIdler interface {
CloseIdleConnections()
}
type unwrapper struct {
http.RoundTripper
closeIdler
}
if _, ok := wrapper.(closeIdler); ok {
return wrapper
}
if c, ok := wrapped.(closeIdler); ok {
return &unwrapper{
RoundTripper: wrapper,
closeIdler: c,
}
}
return wrapper
}
func (s *staticKubeCreds) close() error {
return nil
}
// dynamicCredsClient defines the function signature used by `dynamicCreds`
// to generate and renew short-lived credentials to access the cluster.
type dynamicCredsClient func(ctx context.Context, cluster types.KubeCluster) (cfg *rest.Config, expirationTime time.Time, err error)
// dynamicKubeCreds contains short-lived credentials to access the cluster.
// Unlike `staticKubeCreds`, `dynamicKubeCreds` extracts access credentials using the `client`
// function and renews them whenever they are about to expire.
type dynamicKubeCreds struct {
ctx context.Context
renewTicker clockwork.Ticker
staticCreds *staticKubeCreds
log *slog.Logger
closeC chan struct{}
client dynamicCredsClient
checker servicecfg.ImpersonationPermissionsChecker
clock clockwork.Clock
component KubeServiceType
sync.RWMutex
wg sync.WaitGroup
}
// dynamicCredsConfig contains configuration for dynamicKubeCreds.
type dynamicCredsConfig struct {
kubeCluster types.KubeCluster
log *slog.Logger
client dynamicCredsClient
checker servicecfg.ImpersonationPermissionsChecker
clock clockwork.Clock
initialRenewInterval time.Duration
resourceMatchers []services.ResourceMatcher
component KubeServiceType
}
func (d *dynamicCredsConfig) checkAndSetDefaults() error {
if d.kubeCluster == nil {
return trace.BadParameter("missing kubeCluster")
}
if d.log == nil {
return trace.BadParameter("missing log")
}
if d.client == nil {
return trace.BadParameter("missing client")
}
if d.checker == nil {
return trace.BadParameter("missing checker")
}
if d.clock == nil {
d.clock = clockwork.NewRealClock()
}
if d.initialRenewInterval == 0 {
d.initialRenewInterval = time.Hour
}
return nil
}
// newDynamicKubeCreds creates a new dynamicKubeCreds refresher and starts the
// credentials refresher mechanism to renew them once they are about to expire.
func newDynamicKubeCreds(ctx context.Context, cfg dynamicCredsConfig) (*dynamicKubeCreds, error) {
if err := cfg.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
dyn := &dynamicKubeCreds{
ctx: ctx,
log: cfg.log,
closeC: make(chan struct{}),
client: cfg.client,
renewTicker: cfg.clock.NewTicker(cfg.initialRenewInterval),
checker: cfg.checker,
clock: cfg.clock,
component: cfg.component,
}
if err := dyn.renewClientset(cfg.kubeCluster); err != nil {
return nil, trace.Wrap(err)
}
dyn.wg.Add(1)
go func() {
defer dyn.wg.Done()
for {
select {
case <-dyn.closeC:
return
case <-dyn.renewTicker.Chan():
if err := dyn.renewClientset(cfg.kubeCluster); err != nil {
cfg.log.WarnContext(ctx, "Unable to renew cluster credentials", "cluster", cfg.kubeCluster.GetName(), "error", err)
}
}
}
}()
return dyn, nil
}
func (d *dynamicKubeCreds) getTLSConfig() *tls.Config {
d.RLock()
defer d.RUnlock()
return d.staticCreds.getTLSConfig()
}
func (d *dynamicKubeCreds) getTransportConfig() *transport.Config {
d.RLock()
defer d.RUnlock()
return d.staticCreds.transportConfig
}
func (d *dynamicKubeCreds) getKubeRestConfig() *rest.Config {
d.RLock()
defer d.RUnlock()
return d.staticCreds.clientRestCfg
}
func (d *dynamicKubeCreds) getTargetAddr() string {
d.RLock()
defer d.RUnlock()
return d.staticCreds.targetAddr
}
func (d *dynamicKubeCreds) getKubeClient() kubernetes.Interface {
d.RLock()
defer d.RUnlock()
return d.staticCreds.kubeClient
}
func (d *dynamicKubeCreds) wrapTransport(rt http.RoundTripper) (http.RoundTripper, error) {
d.RLock()
defer d.RUnlock()
return d.staticCreds.wrapTransport(rt)
}
func (d *dynamicKubeCreds) close() error {
close(d.closeC)
d.wg.Wait()
d.renewTicker.Stop()
return nil
}
func (d *dynamicKubeCreds) getTransport() http.RoundTripper {
d.RLock()
defer d.RUnlock()
return d.staticCreds.getTransport()
}
// renewClientset generates the credentials required for accessing the cluster using the client function.
func (d *dynamicKubeCreds) renewClientset(cluster types.KubeCluster) error {
// get auth config
restConfig, exp, err := d.client(d.ctx, cluster)
if err != nil {
return trace.Wrap(err)
}
creds, err := extractKubeCreds(d.ctx, d.component, cluster.GetName(), restConfig, d.log, d.checker)
if err != nil {
return trace.Wrap(err)
}
d.Lock()
defer d.Unlock()
d.staticCreds = creds
// prepares the next renew cycle
if !exp.IsZero() {
reset := exp.Sub(d.clock.Now()) / 2
d.renewTicker.Reset(reset)
}
return nil
}
/*
Copyright 2016 The Kubernetes Authors.
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.
*/
package proxy
import (
"context"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"strconv"
"sync"
"time"
"github.com/gravitational/trace"
"k8s.io/streaming/pkg/httpstream"
spdystream "k8s.io/streaming/pkg/httpstream/spdy"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/utils"
)
// portForwardRequest is a request that specifies port forwarding
type portForwardRequest struct {
podNamespace string
podName string
ports []string
httpRequest *http.Request
httpResponseWriter http.ResponseWriter
onPortForward portForwardCallback
context context.Context
targetDialer httpstream.Dialer
pingPeriod time.Duration
idleTimeout time.Duration
}
func (p portForwardRequest) String() string {
return fmt.Sprintf("port forward %v/%v -> %v", p.podNamespace, p.podName, p.ports)
}
// portForwardCallback is a callback to be called on every port forward request
type portForwardCallback func(addr string, success bool)
// parsePortString parses a port from a given string.
func parsePortString(pString string) (uint16, error) {
port, err := strconv.ParseUint(pString, 10, 16)
if err != nil {
return 0, trace.BadParameter("unable to parse %q as a port: %v", pString, err)
}
if port < 1 {
return 0, trace.BadParameter("port %q must be > 0", pString)
}
return uint16(port), nil
}
// runPortForwardingHTTPStreams upgrades the clients using SPDY protocol.
// It supports multiplexing and HTTP streams and can be used per-request.
func runPortForwardingHTTPStreams(req portForwardRequest) error {
targetConn, _, err := req.targetDialer.Dial(PortForwardProtocolV1Name)
if err != nil {
return trace.Wrap(err)
}
defer targetConn.Close()
_, err = httpstream.Handshake(req.httpRequest, req.httpResponseWriter, []string{PortForwardProtocolV1Name})
if err != nil {
return trace.Wrap(err)
}
streamChan := make(chan httpstream.Stream, 1)
upgrader := spdystream.NewResponseUpgraderWithPings(req.pingPeriod)
conn := upgrader.UpgradeResponse(req.httpResponseWriter, req.httpRequest, httpStreamReceived(req.context, streamChan))
if conn == nil {
return trace.ConnectionProblem(nil, "Unable to upgrade websocket connection")
}
defer conn.Close()
h := &portForwardProxy{
logger: slog.With(
teleport.ComponentKey, teleport.Component(teleport.ComponentProxyKube),
events.RemoteAddr, req.httpRequest.RemoteAddr,
),
portForwardRequest: req,
sourceConn: conn,
streamChan: streamChan,
streamPairs: make(map[string]*httpStreamPair),
streamCreationTimeout: DefaultStreamCreationTimeout,
targetConn: targetConn,
}
defer h.Close()
h.logger.DebugContext(req.context, "Setting port forwarding streaming connection idle timeout", "idle_timeout", req.idleTimeout)
conn.SetIdleTimeout(adjustIdleTimeoutForConn(req.idleTimeout))
h.run()
return nil
}
// httpStreamReceived is the httpstream.NewStreamHandler for port
// forward streams. It checks each stream's port and stream type headers,
// rejecting any streams that with missing or invalid values. Each valid
// stream is sent to the streams channel.
func httpStreamReceived(ctx context.Context, streams chan httpstream.Stream) func(httpstream.Stream, <-chan struct{}) error {
return func(stream httpstream.Stream, replySent <-chan struct{}) error {
// make sure it has a valid port header
portString := stream.Headers().Get(PortHeader)
if len(portString) == 0 {
return trace.BadParameter("%q header is required", PortHeader)
}
_, err := parsePortString(portString)
if err != nil {
return trace.Wrap(err)
}
// make sure it has a valid stream type header
streamType := stream.Headers().Get(StreamType)
if len(streamType) == 0 {
return trace.BadParameter("%q header is required", StreamType)
}
if streamType != StreamTypeError && streamType != StreamTypeData {
return trace.BadParameter("invalid stream type %q", streamType)
}
select {
case streams <- stream:
return nil
case <-ctx.Done():
return trace.BadParameter("request has been canceled")
}
}
}
// portForwardProxy is capable of processing multiple port forward
// requests over a single httpstream.Connection.
type portForwardProxy struct {
logger *slog.Logger
portForwardRequest
sourceConn httpstream.Connection
streamChan chan httpstream.Stream
streamPairsLock sync.RWMutex
streamPairs map[string]*httpStreamPair
streamCreationTimeout time.Duration
targetConn httpstream.Connection
}
func (h *portForwardProxy) Close() error {
if h.sourceConn != nil {
return h.sourceConn.Close()
}
return nil
}
// forwardStreamPair creates a new data and error streams using the same requestID
// received from the client and copies the data between target's data and error and
// client's data and error streams. It blocks until all copy operations complete.
// It does not close the client's data and error streams as they are closed by
// the caller.
func (h *portForwardProxy) forwardStreamPair(p *httpStreamPair, remotePort int64) error {
// create error stream
headers := http.Header{}
port := fmt.Sprintf("%d", remotePort)
headers.Set(StreamType, StreamTypeError)
headers.Set(PortHeader, port)
headers.Set(PortForwardRequestIDHeader, p.requestID)
// read and write from the error stream
targetErrorStream, err := h.targetConn.CreateStream(headers)
h.onPortForward(net.JoinHostPort(h.podName, port), err == nil /* success */)
if err != nil {
err := trace.ConnectionProblem(err, "error creating error stream for port %d", remotePort)
p.sendErr(err)
return err
}
defer func() {
// on stream close, remove the stream from the connection and close it.
h.targetConn.RemoveStreams(targetErrorStream)
targetErrorStream.Close()
}()
wg := &sync.WaitGroup{}
wg.Add(1)
go func() {
defer wg.Done()
// Close the target error stream to indicate no more writes.
if err := targetErrorStream.Close(); err != nil {
h.logger.DebugContext(h.context, "Unable to close target error stream", "error", err)
}
// Enables error propagation from Kube API server to kubectl client.
if _, err := io.Copy(p.errorStream, targetErrorStream); err != nil {
h.logger.DebugContext(h.context, "Unable to proxy portforward error-stream", "error", err)
}
}()
// create data stream
headers.Set(StreamType, StreamTypeData)
targetDataStream, err := h.targetConn.CreateStream(headers)
if err != nil {
err := trace.ConnectionProblem(err, "error creating forwarding stream for port -> %d: %v", remotePort, err)
p.sendErr(err)
return err
}
defer func() {
// on stream close, remove the stream from the connection and close it.
h.targetConn.RemoveStreams(targetDataStream)
targetDataStream.Close()
}()
wg.Add(1)
go func() {
defer wg.Done()
if err := utils.ProxyConn(h.context, p.dataStream, targetDataStream); err != nil {
h.logger.DebugContext(h.context, "Unable to proxy portforward data-stream", "error", err)
}
}()
h.logger.DebugContext(h.context, "Streams have been created, Waiting for copy to complete")
// wait for the copies to complete before returning.
wg.Wait()
h.logger.DebugContext(h.context, "Port forwarding pair completed")
return nil
}
// getStreamPair returns a httpStreamPair for requestID. This creates a
// new pair if one does not yet exist for the requestID. The returned bool is
// true if the pair was created.
func (h *portForwardProxy) getStreamPair(requestID string) (*httpStreamPair, bool) {
h.streamPairsLock.Lock()
defer h.streamPairsLock.Unlock()
if p, ok := h.streamPairs[requestID]; ok {
h.logger.DebugContext(h.context, "Found existing stream pair for request", "request_id", requestID)
return p, false
}
h.logger.DebugContext(h.context, "Creating new stream pair for request", "request_id", requestID)
p := newPortForwardPair(requestID)
h.streamPairs[requestID] = p
return p, true
}
// monitorStreamPair waits for the pair to receive both its error and data
// streams, or for the timeout to expire (whichever happens first), and then
// removes the pair.
func (h *portForwardProxy) monitorStreamPair(p *httpStreamPair) {
timeC := time.NewTimer(h.streamCreationTimeout)
defer timeC.Stop()
select {
case <-timeC.C:
h.logger.ErrorContext(h.context, "Request timed out waiting for streams", "request_id", p.requestID)
case <-p.complete:
h.logger.DebugContext(h.context, "Request successfully received error and data streams", "request_id", p.requestID)
}
h.removeStreamPair(p.requestID)
}
// removeStreamPair removes the stream pair identified by requestID from streamPairs.
func (h *portForwardProxy) removeStreamPair(requestID string) {
h.streamPairsLock.Lock()
defer h.streamPairsLock.Unlock()
pair, ok := h.streamPairs[requestID]
if !ok {
return
}
if h.sourceConn != nil {
// remove the streams from the connection and close them.
h.sourceConn.RemoveStreams(pair.dataStream, pair.errorStream)
}
delete(h.streamPairs, requestID)
}
// requestID returns the request id for stream.
func (h *portForwardProxy) requestID(stream httpstream.Stream) (string, error) {
requestID := stream.Headers().Get(PortForwardRequestIDHeader)
if len(requestID) == 0 {
return "", trace.BadParameter("port forwarding is not supported")
}
return requestID, nil
}
// run is the main loop for the portForwardProxy. It processes new
// streams, invoking portForward for each complete stream pair. The loop exits
// when the httpstream.Connection is closed.
func (h *portForwardProxy) run() {
h.logger.DebugContext(h.context, "Waiting for port forward streams")
var wg sync.WaitGroup
defer wg.Wait()
for {
select {
case <-h.context.Done():
h.logger.DebugContext(h.context, "Context is closing, returning")
return
case <-h.sourceConn.CloseChan():
h.logger.DebugContext(h.context, "Upgraded connection closed")
return
case <-h.targetConn.CloseChan():
h.logger.DebugContext(h.context, "Target connection closed")
return
case stream := <-h.streamChan:
requestID, err := h.requestID(stream)
if err != nil {
h.logger.WarnContext(h.context, "Failed to parse request id", "error", err)
return
}
streamType := stream.Headers().Get(StreamType)
h.logger.DebugContext(h.context, "Received new stream", "request_id", requestID, "stream_type", streamType)
p, created := h.getStreamPair(requestID)
if created {
go h.monitorStreamPair(p)
}
if complete, err := p.add(stream); err != nil {
err := trace.BadParameter("error processing stream for request %s: %v", requestID, err)
p.sendErr(err)
} else if complete {
wg.Add(1)
go func() {
defer wg.Done()
h.portForward(p)
}()
}
}
}
}
// portForward handles the port-forwarding for the given stream pair.
// It closes the pair when it is done.
func (h *portForwardProxy) portForward(p *httpStreamPair) {
defer p.close()
portString := p.dataStream.Headers().Get(PortHeader)
port, _ := strconv.ParseInt(portString, 10, 32)
logger := h.logger.With("request_id", p.requestID, "port", portString)
logger.DebugContext(h.context, "Forwarding port")
if err := h.forwardStreamPair(p, port); err != nil {
logger.DebugContext(h.context, "Error forwarding port", "error", err)
return
}
h.logger.DebugContext(h.context, "Completed forwarding port")
}
// httpStreamPair represents the error and data streams for a port
// forwarding request.
type httpStreamPair struct {
lock sync.Mutex
requestID string
dataStream httpstream.Stream
errorStream httpstream.Stream
complete chan struct{}
}
// newPortForwardPair creates a new httpStreamPair.
func newPortForwardPair(requestID string) *httpStreamPair {
return &httpStreamPair{
requestID: requestID,
complete: make(chan struct{}),
}
}
// add adds the stream to the httpStreamPair. If the pair already
// contains a stream for the new stream's type, an error is returned. add
// returns true if both the data and error streams for this pair have been
// received.
func (p *httpStreamPair) add(stream httpstream.Stream) (bool, error) {
p.lock.Lock()
defer p.lock.Unlock()
switch stream.Headers().Get(StreamType) {
case StreamTypeError:
if p.errorStream != nil {
return false, trace.BadParameter("error stream already assigned")
}
p.errorStream = stream
case StreamTypeData:
if p.dataStream != nil {
return false, trace.BadParameter("data stream already assigned")
}
p.dataStream = stream
}
complete := p.errorStream != nil && p.dataStream != nil
if complete {
close(p.complete)
}
return complete, nil
}
// sendErr writes s to p.errorStream if p.errorStream has been set.
func (p *httpStreamPair) sendErr(err error) {
if err == nil {
return
}
p.lock.Lock()
defer p.lock.Unlock()
if p.errorStream != nil {
fmt.Fprint(p.errorStream, err.Error())
}
}
// close closes the data and error streams for this pair.
func (p *httpStreamPair) close() {
p.lock.Lock()
defer p.lock.Unlock()
if p.dataStream != nil {
p.dataStream.Close()
}
if p.errorStream != nil {
p.errorStream.Close()
}
}
/*
* 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 proxy
import (
"context"
"encoding/binary"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"sync"
gwebsocket "github.com/gorilla/websocket"
"github.com/gravitational/trace"
portforwardconstants "k8s.io/apimachinery/pkg/util/portforward"
"k8s.io/client-go/tools/portforward"
"k8s.io/streaming/pkg/httpstream"
spdystream "k8s.io/streaming/pkg/httpstream/spdy"
"k8s.io/streaming/pkg/httpstream/wsstream"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/utils"
)
const (
// portForwardDataChannel is the prefix for WebSocket data channel.
// data: [portForwardDataChannel, data...]
portForwardDataChannel = iota
// portForwardErrorChannel is the prefix for WebSocket error channel.
// error: [portForwardErrorChannel, data...]
portForwardErrorChannel
)
// runPortForwardingWebSocket handles a request to forward ports to a pod using
// WebSocket protocol. For each port to forward, a pair of "channels" is created
// (DATA (0), ERROR (1)) when the request is upgraded and the associated port is
// written to each channel as unsigned 16 integer. It's required to identify
// which channels belong to each port.
// Due to a protocol limitation, WebSockets do not support multiplexing nor
// concurrent requests.
func runPortForwardingWebSocket(req portForwardRequest) error {
// When dialing to the upstream server (Teleport or Kubernetes API server),
// Teleport uses the SPDY implementation instead of WebSockets.
targetConn, _, err := req.targetDialer.Dial(PortForwardProtocolV1Name)
if err != nil {
return trace.Wrap(err, "error dialing to upstream connection")
}
defer targetConn.Close()
ports, err := extractTargetPortsFromStrings(req.ports)
if err != nil {
return trace.Wrap(err)
}
// One pair of (Data,Error) channels per port.
channels := make([]wsstream.ChannelType, 2*len(ports))
for i := range channels {
channels[i] = wsstream.ReadWriteChannel
}
// Create a stream upgrader with protocol negotiation.
conn := wsstream.NewConn(map[string]wsstream.ChannelProtocolConfig{
"": {
Binary: true,
Channels: channels,
},
v4BinaryWebsocketProtocol: {
Binary: true,
Channels: channels,
},
v4Base64WebsocketProtocol: {
Binary: false,
Channels: channels,
},
})
conn.SetIdleTimeout(adjustIdleTimeoutForConn(req.idleTimeout))
// Upgrade the request and create the virtual streams.
_, streams, err := conn.Open(
req.httpResponseWriter,
req.httpRequest,
)
if err != nil {
return trace.ConnectionProblem(err, "unable to upgrade websocket connection")
}
defer conn.Close()
// Create the websocket stream pairs.
streamPairs := make([]*websocketChannelPair, len(ports))
for i := range ports {
var (
dataStream = streams[2*i+portForwardDataChannel]
errorStream = streams[2*i+portForwardErrorChannel]
port = ports[i]
)
streamPairs[i] = &websocketChannelPair{
port: port,
dataStream: dataStream,
errorStream: errorStream,
// create one requestID per pair so we can forward to multiple ports
// correctly.
// Since websockets do no support multiplexing, it's ok to use a single
// request per port since users cannot send concurrent requests to
// Kubernetes API server.
// Although users can connect via Websocket, Teleport connection between
// its components or Kubernetes API server is done using SPDY client
// which requires request_id.
requestID: fmt.Sprintf("%d", port),
podName: req.podName,
}
portBytes := make([]byte, 2)
binary.LittleEndian.PutUint16(portBytes, port)
// Protocol requires sending the port to identify which channels belong to
// each port.
if _, err := dataStream.Write(portBytes); err != nil {
return trace.Wrap(err)
}
if _, err := errorStream.Write(portBytes); err != nil {
return trace.Wrap(err)
}
}
h := &websocketPortforwardHandler{
conn: conn,
streamPairs: streamPairs,
podName: req.podName,
targetConn: targetConn,
onPortForward: req.onPortForward,
logger: slog.With(
teleport.ComponentKey, teleport.Component(teleport.ComponentProxyKube),
events.RemoteAddr, req.httpRequest.RemoteAddr,
),
context: req.context,
}
// run the portforward request until termination.
h.run()
return nil
}
// extractTargetPortsFromStrings extracts the desired ports from the request
// query parameters.
func extractTargetPortsFromStrings(portsStrings []string) ([]uint16, error) {
if len(portsStrings) == 0 {
return nil, trace.BadParameter("query parameter %q is required", PortHeader)
}
ports := make([]uint16, 0, len(portsStrings))
for _, portString := range portsStrings {
if len(portString) == 0 {
return nil, trace.BadParameter("query parameter %q cannot be empty", PortHeader)
}
for p := range strings.SplitSeq(portString, ",") {
port, err := parsePortString(p)
if err != nil {
return nil, trace.Wrap(err)
}
ports = append(ports, port)
}
}
return ports, nil
}
// websocketChannelPair represents the error and data streams for a single
// port.
type websocketChannelPair struct {
port uint16
podName string
requestID string
dataStream io.ReadWriteCloser
errorStream io.ReadWriteCloser
}
func (w *websocketChannelPair) close() {
w.dataStream.Close()
w.errorStream.Close()
}
func (w *websocketChannelPair) sendErr(err error) {
if err == nil {
return
}
fmt.Fprintf(w.errorStream, "error forwarding port %d to pod %s: %v", w.port, w.podName, err)
}
// websocketPortforwardHandler is capable of processing a single port forward
// request over a websocket connection
type websocketPortforwardHandler struct {
conn *wsstream.Conn
streamPairs []*websocketChannelPair
podName string
targetConn httpstream.Connection
onPortForward portForwardCallback
logger *slog.Logger
context context.Context
}
// run invokes the targetConn SPDY connection and copies the client data into
// the targetConn and the responses into the targetConn data stream.
// If any error occurs, stream is closed an the error is sent via errorStream.
func (h *websocketPortforwardHandler) run() {
wg := sync.WaitGroup{}
wg.Add(len(h.streamPairs))
for _, pair := range h.streamPairs {
p := pair
go func() {
defer wg.Done()
h.portForward(p)
}()
}
wg.Wait()
}
// portForward copies the client and upstream streams.
func (h *websocketPortforwardHandler) portForward(p *websocketChannelPair) {
logger := h.logger.With("request_id", p.requestID, "port", p.port)
logger.DebugContext(h.context, "Forwarding port")
h.forwardStreamPair(p)
logger.DebugContext(h.context, "Completed forwarding port")
}
func (h *websocketPortforwardHandler) forwardStreamPair(p *websocketChannelPair) {
// create error stream
headers := http.Header{}
headers.Set(StreamType, StreamTypeError)
headers.Set(PortHeader, fmt.Sprint(p.port))
headers.Set(PortForwardRequestIDHeader, p.requestID)
// read and write from the error stream
targetErrorStream, err := h.targetConn.CreateStream(headers)
h.onPortForward(fmt.Sprintf("%v:%v", h.podName, p.port), err == nil /* success */)
if err != nil {
p.sendErr(err)
return
}
defer func() {
// on stream close, remove the stream from the connection and close it.
h.targetConn.RemoveStreams(targetErrorStream)
targetErrorStream.Close()
}()
wg := &sync.WaitGroup{}
wg.Add(1)
go func() {
defer wg.Done()
// Close the target error stream to indicate no more writes.
if err := targetErrorStream.Close(); err != nil {
h.logger.DebugContext(h.context, "Unable to close target error stream", "error", err)
}
// Enables error propagation from Kube API server to kubectl client.
if _, err := io.Copy(p.errorStream, targetErrorStream); err != nil {
h.logger.DebugContext(h.context, "Unable to proxy portforward error-stream", "error", err)
}
}()
// create data stream
headers.Set(StreamType, StreamTypeData)
targetDataStream, err := h.targetConn.CreateStream(headers)
if err != nil {
p.sendErr(err)
p.close()
wg.Wait()
return
}
defer func() {
// on stream close, remove the stream from the connection and close it.
h.targetConn.RemoveStreams(targetDataStream)
targetDataStream.Close()
}()
wg.Add(1)
go func() {
defer wg.Done()
if err := utils.ProxyConn(h.context, p.dataStream, targetDataStream); err != nil {
h.logger.DebugContext(h.context, "Unable to proxy portforward data-stream", "error", err)
}
}()
h.logger.DebugContext(h.context, "Streams have been created, Waiting for copy to complete")
// Wait until every goroutine exits.
wg.Wait()
h.logger.DebugContext(h.context, "Port forwarding pair completed")
}
// runPortForwardingTunneledHTTPStreams handles a port-forwarding request that uses SPDY protocol
// over WebSockets.
func runPortForwardingTunneledHTTPStreams(req portForwardRequest) error {
targetConn, _, err := req.targetDialer.Dial(PortForwardProtocolV1Name)
if err != nil {
return trace.Wrap(err)
}
defer targetConn.Close()
// Try to upgrade the websocket connection.
// Beyond this point, we don't need to write errors to the response.
upgrader := gwebsocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
Subprotocols: []string{portforwardconstants.WebsocketsSPDYTunnelingPortForwardV1},
}
conn, err := upgrader.Upgrade(req.httpResponseWriter, req.httpRequest, nil)
if err != nil {
return trace.Wrap(err)
}
tunneledConn := portforward.NewTunnelingConnection("server", conn)
streamChan := make(chan httpstream.Stream, 1)
spdyConn, err := spdystream.NewServerConnectionWithPings(
tunneledConn,
httpStreamReceived(req.context, streamChan),
req.pingPeriod,
)
if err != nil {
return trace.Wrap(err)
}
if conn == nil {
return trace.ConnectionProblem(nil, "Unable to upgrade websocket connection")
}
defer conn.Close()
h := &portForwardProxy{
logger: slog.With(
teleport.ComponentKey, teleport.Component(teleport.ComponentProxyKube),
events.RemoteAddr, req.httpRequest.RemoteAddr,
),
portForwardRequest: req,
sourceConn: spdyConn,
streamChan: streamChan,
streamPairs: make(map[string]*httpStreamPair),
streamCreationTimeout: DefaultStreamCreationTimeout,
targetConn: targetConn,
}
defer h.Close()
h.logger.DebugContext(context.Background(), "Setting port forwarding streaming connection idle timeout to", "idle_timeout", req.idleTimeout)
spdyConn.SetIdleTimeout(adjustIdleTimeoutForConn(req.idleTimeout))
h.run()
return nil
}
/*
* 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 proxy
import (
"net/http"
"github.com/prometheus/client_golang/prometheus"
"github.com/prometheus/client_golang/prometheus/promhttp"
"go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp"
"github.com/gravitational/teleport"
tracehttp "github.com/gravitational/teleport/api/observability/tracing/http"
"github.com/gravitational/teleport/lib/observability/metrics"
)
const (
// kubernetesSubsystem is used to prefix Prometheus metrics for this
// subsystem.
// See https://prometheus.io/docs/practices/naming/#subsystem-name
kubernetesSubsystem = "kubernetes"
)
func init() {
metrics.RegisterPrometheusCollectors(
clientRequestCounter,
clientTLSLatencyVec,
clientRequestDurationHistVec,
clientInFlightGauge,
clietGotConnLatencyVec,
clientFirstByteLatencyVec,
serverInFlightGauge,
serverRequestCounter,
serverRequestDurationHist,
serverResponseSizeHist,
execSessionsInFlightGauge,
execSessionsRequestCounter,
portforwardSessionsInFlightGauge,
portforwardRequestCounter,
joinSessionsInFlightGauge,
joinSessionsRequestCounter,
)
}
// The following section defines Prometheus metrics for the clients used by
// Teleport proxy to connect to the Teleport Kubernetes service and by the
// Teleport Kubernetes service to connect to the Kubernetes cluster.
var (
clientInFlightGauge = prometheus.NewGaugeVec(
prometheus.GaugeOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "client_in_flight_requests",
Help: "In-flight requests waiting for the upstream response.",
},
[]string{"component"},
)
clientRequestCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "client_requests_total",
Help: "Total number of requests sent to the upstream teleport proxy, kube_service or Kubernetes Cluster servers.",
},
[]string{"component", "code", "method"},
)
clientTLSLatencyVec = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "client_tls_duration_seconds",
Help: "Latency distribution of TLS handshakes.",
Buckets: prometheus.DefBuckets,
},
[]string{"component", "event"},
)
clietGotConnLatencyVec = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "client_got_conn_duration_seconds",
Help: "A histogram of latency to dial to the upstream server.",
Buckets: prometheus.DefBuckets,
},
[]string{"component"},
)
clientFirstByteLatencyVec = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "client_first_byte_response_duration_seconds",
Help: "Teleport Kubernetes Service | Latency distribution of time to receive the first response byte from the upstream server.",
Buckets: prometheus.DefBuckets,
},
[]string{"component"},
)
clientRequestDurationHistVec = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "client_request_duration_seconds",
Help: "Latency distribution of the upstream request time.",
Buckets: prometheus.DefBuckets,
},
[]string{"component"},
)
)
// instrumentedRoundtripper instruments the provided RoundTripper with
// Prometheus metrics and OpenTelemetry tracing.
func instrumentedRoundtripper(component string, tr http.RoundTripper) http.RoundTripper {
// Define functions for the available httptrace.ClientTrace hook
// functions that we want to instrument.
httpTrace := &promhttp.InstrumentTrace{
GotConn: func(t float64) {
clietGotConnLatencyVec.WithLabelValues(component).Observe(t)
},
GotFirstResponseByte: func(t float64) {
clientFirstByteLatencyVec.WithLabelValues(component).Observe(t)
},
TLSHandshakeStart: func(t float64) {
clientTLSLatencyVec.WithLabelValues(component, "tls_handshake_start").Observe(t)
},
TLSHandshakeDone: func(t float64) {
clientTLSLatencyVec.WithLabelValues(component, "tls_handshake_done").Observe(t)
},
}
curryWith := prometheus.Labels{"component": component}
return tracehttp.NewTransportWithInner(
promhttp.InstrumentRoundTripperInFlight(
clientInFlightGauge.WithLabelValues(component),
promhttp.InstrumentRoundTripperCounter(
clientRequestCounter.MustCurryWith(curryWith),
promhttp.InstrumentRoundTripperTrace(
httpTrace,
promhttp.InstrumentRoundTripperDuration(clientRequestDurationHistVec.MustCurryWith(curryWith), tr),
),
),
),
// Pass the original RoundTripper to the inner transport so that it can
// be used to close idle connections because promhttp roundtrippers don't
// implement CloseIdleConnections.
tr,
)
}
// The following section defines Prometheus metrics for the HTTP server used by
// the Teleport Kubernetes Proxy and the Teleport Kubernetes service.
var (
serverInFlightGauge = prometheus.NewGaugeVec(
prometheus.GaugeOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_in_flight_requests",
Help: "In-flight requests currently handled by the server.",
},
[]string{"component"},
)
serverRequestCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_api_requests_total",
Help: "Total number of requests handled by the server.",
},
[]string{"component", "code", "method"},
)
serverRequestDurationHist = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_request_duration_seconds",
Help: "Latency distribution of the total request time.",
Buckets: prometheus.DefBuckets,
},
[]string{"component", "method"},
)
serverResponseSizeHist = prometheus.NewHistogramVec(
prometheus.HistogramOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_response_size_bytes",
Help: "Distribution of the response size.",
// The following exponential buckets are equivalent to the following:
// [50B 150B 450B 1.32KB 3.96KB 11.87KB 35.6KB 106.79KB 320.36KB 961.08KB 2.82MB 8.45MB 25.34MB]
Buckets: prometheus.ExponentialBuckets(50, 3, 13),
},
[]string{"component"},
)
execSessionsInFlightGauge = prometheus.NewGaugeVec(
prometheus.GaugeOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_exec_in_flight_sessions",
Help: "Number of active kubectl exec sessions.",
},
[]string{"component"},
)
execSessionsRequestCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_exec_sessions_total",
Help: "Total number of kubectl exec sessions. ",
},
[]string{"component"},
)
portforwardSessionsInFlightGauge = prometheus.NewGaugeVec(
prometheus.GaugeOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_portforward_in_flight_sessions",
Help: " Number of active kubectl portforward sessions.",
},
[]string{"component"},
)
portforwardRequestCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_portforward_sessions_total",
Help: "Number of active kubectl portforward sessions.",
},
[]string{"component"},
)
joinSessionsInFlightGauge = prometheus.NewGaugeVec(
prometheus.GaugeOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_join_in_flight_sessions",
Help: "Number of active joining sessions,",
},
[]string{"component"},
)
joinSessionsRequestCounter = prometheus.NewCounterVec(
prometheus.CounterOpts{
Namespace: teleport.MetricNamespace,
Subsystem: kubernetesSubsystem,
Name: "server_join_sessions_total",
Help: "Total number of joining sessions.",
},
[]string{"component"},
)
)
// instrumentHTTPHandler instruments the HTTP handler with OpenTelemetry and
// Prometheus metrics.
func instrumentHTTPHandler(component string, handler http.Handler) http.Handler {
return otelhttp.NewHandler(
instrumentHTTPHandlerWithPrometheus(component, handler),
component,
otelhttp.WithMessageEvents(otelhttp.ReadEvents, otelhttp.WriteEvents),
)
}
// instrumentHTTPHandlerWithPrometheus instruments the HTTP handler with
// Prometheus metrics.
func instrumentHTTPHandlerWithPrometheus(component string, handler http.Handler) http.Handler {
curryWith := prometheus.Labels{"component": component}
return promhttp.InstrumentHandlerInFlight(
serverInFlightGauge.WithLabelValues(component),
promhttp.InstrumentHandlerDuration(
serverRequestDurationHist.MustCurryWith(
curryWith,
),
promhttp.InstrumentHandlerCounter(
serverRequestCounter.MustCurryWith(
curryWith,
),
promhttp.InstrumentHandlerResponseSize(
serverResponseSizeHist.MustCurryWith(
curryWith,
),
handler,
),
),
),
)
}
/*
Copyright 2016 The Kubernetes Authors.
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.
*/
package proxy
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"strings"
"sync"
"time"
"github.com/gravitational/trace"
apierrors "k8s.io/apimachinery/pkg/api/errors"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime"
remotecommandconsts "k8s.io/apimachinery/pkg/util/remotecommand"
"k8s.io/client-go/tools/remotecommand"
utilexec "k8s.io/client-go/util/exec"
"k8s.io/streaming/pkg/httpstream"
spdystream "k8s.io/streaming/pkg/httpstream/spdy"
"k8s.io/streaming/pkg/httpstream/wsstream"
apievents "github.com/gravitational/teleport/api/types/events"
)
// remoteCommandRequest is a request to execute a remote command
type remoteCommandRequest struct {
podNamespace string
podName string
containerName string
cmd []string
stdin bool
stdout bool
stderr bool
tty bool
httpRequest *http.Request
httpResponseWriter http.ResponseWriter
onResize resizeCallback
context context.Context
pingPeriod time.Duration
idleTimeout time.Duration
}
func (req remoteCommandRequest) eventPodMeta(ctx context.Context, creds kubeCreds) apievents.KubernetesPodMetadata {
meta := apievents.KubernetesPodMetadata{
KubernetesPodName: req.podName,
KubernetesPodNamespace: req.podNamespace,
KubernetesContainerName: req.containerName,
}
if creds == nil || creds.getKubeClient() == nil {
return meta
}
// Optionally, try to get more info about the pod.
//
// This can fail if a user has set tight RBAC rules for teleport. Failure
// here shouldn't prevent a session from starting.
pod, err := creds.getKubeClient().CoreV1().Pods(req.podNamespace).Get(ctx, req.podName, metav1.GetOptions{})
if err != nil {
slog.DebugContext(ctx, "Failed fetching pod from kubernetes API; skipping additional metadata on the audit event", "error", err)
return meta
}
meta.KubernetesNodeName = pod.Spec.NodeName
// If a container name was provided, find its image name.
if req.containerName != "" {
for _, c := range pod.Spec.Containers {
if c.Name == req.containerName {
meta.KubernetesContainerImage = c.Image
break
}
}
}
// If no container name was provided, use the default one.
if req.containerName == "" && len(pod.Spec.Containers) > 0 {
meta.KubernetesContainerName = pod.Spec.Containers[0].Name
meta.KubernetesContainerImage = pod.Spec.Containers[0].Image
}
return meta
}
func upgradeRequestToRemoteCommandProxy(req remoteCommandRequest, exec func(*remoteCommandProxy) error) (any, error) {
var (
proxy *remoteCommandProxy
err error
)
if wsstream.IsWebSocketRequest(req.httpRequest) {
proxy, err = createWebSocketStreams(req)
} else {
proxy, err = createSPDYStreams(req)
}
if err != nil {
return nil, trace.Wrap(err)
}
defer proxy.Close()
if proxy.resizeStream != nil {
proxy.resizeQueue = newTermQueue(req.context, req.onResize)
go proxy.resizeQueue.handleResizeEvents(proxy.resizeStream)
}
err = exec(proxy)
if !isRelevantWebsocketError(err) {
err = nil
}
if err := proxy.sendStatus(err); err != nil {
slog.WarnContext(req.context, "Failed to send status", "error", err)
}
// return rsp=nil, err=nil to indicate that the request has been handled
// by the hijacked connection. If we return an error, the request will be
// considered unhandled and the middleware will try to write the error
// or response into the hicjacked connection, which will fail.
return nil /* rsp */, nil /* err */
}
func createSPDYStreams(req remoteCommandRequest) (*remoteCommandProxy, error) {
protocol, err := httpstream.Handshake(req.httpRequest, req.httpResponseWriter, []string{StreamProtocolV4Name})
if err != nil {
return nil, trace.Wrap(err)
}
streamCh := make(chan streamAndReply)
ctx, cancel := context.WithCancel(req.context)
defer cancel()
upgrader := spdystream.NewResponseUpgraderWithPings(req.pingPeriod)
conn := upgrader.UpgradeResponse(req.httpResponseWriter, req.httpRequest, func(stream httpstream.Stream, replySent <-chan struct{}) error {
select {
case streamCh <- streamAndReply{Stream: stream, replySent: replySent}:
return nil
case <-ctx.Done():
return trace.BadParameter("request has been canceled")
}
})
// from this point on, we can no longer call methods on response
if conn == nil {
// The upgrader is responsible for notifying the client of any errors that
// occurred during upgrading. All we can do is return here at this point
// if we weren't successful in upgrading.
return nil, trace.ConnectionProblem(trace.BadParameter("missing connection"), "missing connection")
}
conn.SetIdleTimeout(adjustIdleTimeoutForConn(req.idleTimeout))
var handler protocolHandler
switch protocol {
case "":
slog.WarnContext(ctx, "Client did not request protocol negotiation")
fallthrough
case StreamProtocolV4Name:
slog.InfoContext(ctx, "Negotiated protocol", "protocol", protocol)
handler = &v4ProtocolHandler{}
default:
err = trace.BadParameter("protocol %v is not supported. upgrade the client", protocol)
return nil, trace.NewAggregate(err, conn.Close())
}
// count the streams client asked for, starting with 1
expectedStreams := 1
if req.stdin {
expectedStreams++
}
if req.stdout {
expectedStreams++
}
if req.stderr {
expectedStreams++
}
if req.tty && handler.supportsTerminalResizing() {
expectedStreams++
}
expired := time.NewTimer(DefaultStreamCreationTimeout)
defer expired.Stop()
proxy, err := handler.waitForStreams(ctx, streamCh, expectedStreams, expired.C)
if err != nil {
return nil, trace.NewAggregate(err, conn.Close())
}
proxy.conn = conn
proxy.tty = req.tty
return proxy, nil
}
// remoteCommandProxy contains the connection and streams used when
// forwarding an attach or execute session into a container.
type remoteCommandProxy struct {
conn io.Closer
stdinStream io.ReadCloser
stdoutStream io.WriteCloser
stderrStream io.WriteCloser
writeStatus func(status *apierrors.StatusError) error
resizeStream io.ReadCloser
tty bool
resizeQueue *termQueue
}
func (s *remoteCommandProxy) Close() error {
if s.conn != nil {
return s.conn.Close()
}
// if resize queue is available release its goroutines to prevent stream leaks.
if s.resizeQueue != nil {
s.resizeQueue.Close()
}
return nil
}
func (s *remoteCommandProxy) options() remotecommand.StreamOptions {
opts := remotecommand.StreamOptions{
Stdout: s.stdoutStream,
Stdin: s.stdinStream,
Stderr: s.stderrStream,
Tty: s.tty,
}
// done to prevent this problem: https://golang.org/doc/faq#nil_error
if s.resizeQueue != nil {
opts.TerminalSizeQueue = s.resizeQueue
}
return opts
}
func (s *remoteCommandProxy) sendStatus(err error) error {
if err == nil {
return s.writeStatus(&apierrors.StatusError{ErrStatus: metav1.Status{
Status: metav1.StatusSuccess,
}})
}
var statusErr *apierrors.StatusError
if errors.As(err, &statusErr) {
return s.writeStatus(statusErr)
}
var exitErr utilexec.ExitError
if errors.As(err, &exitErr) && exitErr.Exited() {
rc := exitErr.ExitStatus()
return s.writeStatus(&apierrors.StatusError{ErrStatus: metav1.Status{
Status: metav1.StatusFailure,
Reason: remotecommandconsts.NonZeroExitCodeReason,
Details: &metav1.StatusDetails{
Causes: []metav1.StatusCause{
{
Type: remotecommandconsts.ExitCodeCauseType,
Message: fmt.Sprintf("%d", rc),
},
},
},
Message: fmt.Sprintf("command terminated with non-zero exit code: %v", exitErr),
}})
}
// kubernetes client-go errorDecoderV4 parses the metav1.Status and returns the `fmt.Errorf(status.Message)` for every case except
// errors with reason = NonZeroExitCodeReason for which it returns an exec.CodeExitError.
// This means when forwarding an exec request to a remote cluster using the `Forwarder.remoteExec` function we only have access
// to the status.Message. This happens because the error is sent after the connection was upgraded to a bidirectional stream.
// This hack is here to recreate the forbidden message and return it back to the user terminal
if strings.Contains(err.Error(), "is forbidden:") {
return s.writeStatus(&apierrors.StatusError{
ErrStatus: metav1.Status{
Status: metav1.StatusFailure,
Code: http.StatusForbidden,
Reason: metav1.StatusReasonForbidden,
Message: formatExecForbiddenErrorMessage(err),
},
})
} else if isSessionTerminatedError(err) {
return s.writeStatus(sessionTerminatedByModeratorErr)
}
err = trace.BadParameter("error executing command in container: %v", err)
return s.writeStatus(apierrors.NewInternalError(err))
}
// formatExecForbiddenErrorMessage formats the error message for the forbidden error
// when trying to exec into a pod in Kubernetes 1.30.
func formatExecForbiddenErrorMessage(err error) string {
message := err.Error()
// forbiddenGetResource is the error message that is returned when the user is forbidden to exec into a pod.
// This error message is returned when the user does not have the necessary RBAC rules to exec into a pod.
const forbiddenGetResource = "cannot get resource \"pods/exec\" in API group"
// Kubernetes 1.30 switched to a new exec API that uses a different protocol.
// Previously, the exec API used SPDY. The new exec API uses websockets.
// SPDY allowed the client to send the request as GET or POST, but websockets
// only allow GET per definition.
// This means that most clients that used kubectl version 1.29 or older can
// suddenly get a forbidden error when trying to exec into a pod in Kubernetes 1.30
// because the RBAC rules are not allowing the user to access the pods/exec resource
// using the GET verb.
// This error message is a hint to the user that they need to update their RBAC rules.
if strings.Contains(message, forbiddenGetResource) {
message += kubernetes130BreakingChangeHint
}
return message
}
// streamAndReply holds both a Stream and a channel that is closed when the stream's reply frame is
// enqueued. Consumers can wait for replySent to be closed prior to proceeding, to ensure that the
// replyFrame is enqueued before the connection's goaway frame is sent (e.g. if a stream was
// received and right after, the connection gets closed).
type streamAndReply struct {
httpstream.Stream
replySent <-chan struct{}
}
func newTermQueue(parentContext context.Context, onResize resizeCallback) *termQueue {
ctx, cancel := context.WithCancel(parentContext)
return &termQueue{
ch: make(chan remotecommand.TerminalSize),
cancel: cancel,
done: ctx,
onResize: onResize,
}
}
type resizeCallback func(remotecommand.TerminalSize)
type termQueue struct {
ch chan remotecommand.TerminalSize
cancel context.CancelFunc
done context.Context
onResize resizeCallback
}
func (t *termQueue) Next() *remotecommand.TerminalSize {
select {
case size := <-t.ch:
t.onResize(size)
return &size
case <-t.done.Done():
return nil
}
}
func (t *termQueue) Close() {
t.cancel()
}
func (t *termQueue) handleResizeEvents(stream io.Reader) {
decoder := json.NewDecoder(stream)
for {
size := remotecommand.TerminalSize{}
if err := decoder.Decode(&size); err != nil {
if !errors.Is(err, io.EOF) {
slog.WarnContext(t.done, "Failed to decode resize event", "error", err)
}
t.cancel()
return
}
select {
case t.ch <- size:
case <-t.done.Done():
return
}
}
}
type protocolHandler interface {
// waitForStreams waits for the expected streams or a timeout, returning a
// remoteCommandContext if all the streams were received, or an error if not.
waitForStreams(ctx context.Context, streams <-chan streamAndReply, expectedStreams int, expired <-chan time.Time) (*remoteCommandProxy, error)
// supportsTerminalResizing returns true if the protocol handler supports terminal resizing
supportsTerminalResizing() bool
}
// v4ProtocolHandler implements the V4 protocol version for streaming command execution. It only differs
// in from v3 in the error stream format using an json-marshaled metav1.Status which carries
// the process' exit code.
type v4ProtocolHandler struct{}
func (*v4ProtocolHandler) waitForStreams(connContext context.Context, streams <-chan streamAndReply, expectedStreams int, expired <-chan time.Time) (*remoteCommandProxy, error) {
remoteProxy := &remoteCommandProxy{}
receivedStreams := 0
replyChan := make(chan struct{})
seen := make(map[string]bool, 5)
stopCtx, cancel := context.WithCancel(connContext)
defer cancel()
WaitForStreams:
for {
select {
case stream := <-streams:
streamType := stream.Headers().Get(StreamType)
if seen[streamType] {
return nil, trace.BadParameter("client opened duplicate %q stream", streamType)
}
switch streamType {
case StreamTypeError:
remoteProxy.writeStatus = v4WriteStatusFunc(stream)
seen[streamType] = true
go waitStreamReply(stopCtx, stream.replySent, replyChan)
case StreamTypeStdin:
remoteProxy.stdinStream = stream
seen[streamType] = true
go waitStreamReply(stopCtx, stream.replySent, replyChan)
case StreamTypeStdout:
remoteProxy.stdoutStream = stream
seen[streamType] = true
go waitStreamReply(stopCtx, stream.replySent, replyChan)
case StreamTypeStderr:
remoteProxy.stderrStream = stream
seen[streamType] = true
go waitStreamReply(stopCtx, stream.replySent, replyChan)
case StreamTypeResize:
remoteProxy.resizeStream = stream
seen[streamType] = true
go waitStreamReply(stopCtx, stream.replySent, replyChan)
default:
slog.WarnContext(stopCtx, "Ignoring unexpected stream type", "stream_type", streamType)
}
case <-replyChan:
receivedStreams++
if receivedStreams == expectedStreams {
break WaitForStreams
}
case <-expired:
return nil, trace.BadParameter("timed out waiting for client to create streams")
case <-connContext.Done():
return nil, trace.BadParameter("connection has dropped, exiting")
}
}
return remoteProxy, nil
}
// supportsTerminalResizing returns true because v4ProtocolHandler supports it
func (*v4ProtocolHandler) supportsTerminalResizing() bool { return true }
// waitStreamReply waits until either replySent or stop is closed. If replySent is closed, it sends
// an empty struct to the notify channel.
func waitStreamReply(ctx context.Context, replySent <-chan struct{}, notify chan<- struct{}) {
select {
case <-replySent:
select {
case notify <- struct{}{}:
case <-ctx.Done():
}
case <-ctx.Done():
}
}
// v4WriteStatusFunc returns a WriteStatusFunc that marshals a given api Status
// as json in the error channel.
func v4WriteStatusFunc(stream io.Writer) func(status *apierrors.StatusError) error {
return writeStatusOnceFunc(func(status *apierrors.StatusError) error {
st := status.Status()
data, err := runtime.Encode(globalKubeCodecs.LegacyCodec(), &st)
if err != nil {
return trace.Wrap(err)
}
_, err = stream.Write(data)
return err
})
}
func v1WriteStatusFunc(stream io.Writer) func(status *apierrors.StatusError) error {
return writeStatusOnceFunc(func(status *apierrors.StatusError) error {
if status.Status().Status == metav1.StatusSuccess {
return nil // send error messages
}
_, err := stream.Write([]byte(status.Error()))
return err
})
}
// writeStatusOnceFunc returns a function that only calls f once, and returns the result of the first call.
func writeStatusOnceFunc(f func(status *apierrors.StatusError) error) func(status *apierrors.StatusError) error {
var once sync.Once
var err error
return func(status *apierrors.StatusError) error {
once.Do(func() {
err = f(status)
})
return trace.Wrap(err)
}
}
/*
Copyright 2016 The Kubernetes Authors.
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.
*/
// Origin: https://github.com/kubernetes/kubernetes/blob/d5fdf3135e7c99e5f81e67986ae930f6a2ffb047/pkg/kubelet/cri/streaming/remotecommand/websocket.go
package proxy
import (
"time"
"github.com/go-logr/logr"
"github.com/gravitational/trace"
"k8s.io/apimachinery/pkg/util/remotecommand"
"k8s.io/apiserver/pkg/endpoints/responsewriter"
"k8s.io/klog/v2"
"k8s.io/streaming/pkg/httpstream/wsstream"
)
const (
preV4BinaryWebsocketProtocol = wsstream.ChannelWebSocketProtocol
preV4Base64WebsocketProtocol = wsstream.Base64ChannelWebSocketProtocol
v4BinaryWebsocketProtocol = "v4." + wsstream.ChannelWebSocketProtocol
v4Base64WebsocketProtocol = "v4." + wsstream.Base64ChannelWebSocketProtocol
v5BinaryWebsocketProtocol = remotecommand.StreamProtocolV5Name
)
func init() {
// Replace the default logger from Kubernetes klog package with one that does not log anything.
// This is required to suppress log messages from wsstream when forcing the connection to close.
// Error logs are emitted because `wsstream` does not properly close websocket connections -
// instead of closing only the server side it closes the full connection while the server is
// still waiting for the client to close it.
// Examples of logs emitted by bad behavior are:
// - Use of closed network connection
// - Error on socket receive: read tcp 192.168.1.236:3027->192.168.1.236:57842: use of closed
// network connection
// Go init running order guarantees that the klog package is initialized before this package.
klog.SetLoggerWithOptions(logr.Discard())
}
// createChannels returns the standard channel types for a shell connection (STDIN 0, STDOUT 1, STDERR 2)
// along with the approximate duplex value. It also creates the error (3) and resize (4) channels.
func createChannels(req remoteCommandRequest) []wsstream.ChannelType {
// open the requested channels, and always open the error channel
channels := make([]wsstream.ChannelType, 5)
channels[remotecommand.StreamStdIn] = readChannel(req.stdin)
channels[remotecommand.StreamStdOut] = writeChannel(req.stdout)
channels[remotecommand.StreamStdErr] = writeChannel(req.stderr)
channels[remotecommand.StreamErr] = wsstream.WriteChannel
channels[remotecommand.StreamResize] = wsstream.ReadChannel
return channels
}
// readChannel returns wsstream.ReadChannel if real is true, or wsstream.IgnoreChannel.
func readChannel(real bool) wsstream.ChannelType {
if real {
return wsstream.ReadChannel
}
return wsstream.IgnoreChannel
}
// writeChannel returns wsstream.WriteChannel if real is true, or wsstream.IgnoreChannel.
func writeChannel(real bool) wsstream.ChannelType {
if real {
return wsstream.WriteChannel
}
return wsstream.IgnoreChannel
}
// createWebSocketStreams returns a context containing the websocket connection and
// streams needed to perform an exec or an attach.
func createWebSocketStreams(req remoteCommandRequest) (*remoteCommandProxy, error) {
channels := createChannels(req)
conn := wsstream.NewConn(map[string]wsstream.ChannelProtocolConfig{
"": {
Binary: true,
Channels: channels,
},
preV4BinaryWebsocketProtocol: {
Binary: true,
Channels: channels,
},
preV4Base64WebsocketProtocol: {
Binary: false,
Channels: channels,
},
v4BinaryWebsocketProtocol: {
Binary: true,
Channels: channels,
},
v4Base64WebsocketProtocol: {
Binary: false,
Channels: channels,
},
v5BinaryWebsocketProtocol: {
Binary: true,
Channels: channels,
},
})
conn.SetIdleTimeout(adjustIdleTimeoutForConn(req.idleTimeout))
negotiatedProtocol, streams, err := conn.Open(
responsewriter.GetOriginal(req.httpResponseWriter),
req.httpRequest,
)
if err != nil {
return nil, trace.Wrap(err, "unable to upgrade websocket connection")
}
// Send an empty message to the lowest writable channel to notify the client the connection is established
switch {
case req.stdout:
streams[remotecommand.StreamStdOut].Write([]byte{})
case req.stderr:
streams[remotecommand.StreamStdErr].Write([]byte{})
default:
streams[remotecommand.StreamStdErr].Write([]byte{})
}
proxy := &remoteCommandProxy{
conn: conn,
stdinStream: streams[remotecommand.StreamStdIn],
stdoutStream: streams[remotecommand.StreamStdOut],
stderrStream: streams[remotecommand.StreamStdErr],
tty: req.tty,
resizeStream: streams[remotecommand.StreamResize],
}
// When stdin, stdout or stderr are not enabled, websocket creates a io.Pipe
// for them so they are not nil.
// Since we need to forward to another k8s server (Teleport or real k8s API),
// we must disabled the readers, otherwise the SPDY executor will wait for
// read/write into the streams and will hang.
if !req.stdin {
proxy.stdinStream = nil
}
if !req.stdout {
proxy.stdoutStream = nil
}
if !req.stderr {
proxy.stderrStream = nil
}
switch negotiatedProtocol {
case v5BinaryWebsocketProtocol, v4BinaryWebsocketProtocol, v4Base64WebsocketProtocol:
proxy.writeStatus = v4WriteStatusFunc(streams[remotecommand.StreamErr])
default:
proxy.writeStatus = v1WriteStatusFunc(streams[remotecommand.StreamErr])
}
return proxy, nil
}
// adjustIdleTimeoutForConn adjusts the idle timeout for the connection
// to be 5 seconds longer than the requested idle timeout.
// This is done to prevent the connection from being closed by the server
// before the connection monitor has a chance to close it and write the
// status code.
// If the idle timeout is 0, this function returns 0 because it means the
// connection will never be closed by the server due to idleness.
func adjustIdleTimeoutForConn(idleTimeout time.Duration) time.Duration {
// If the idle timeout is 0, we don't need to adjust it because it
// means the connection will never be closed by the server due to idleness.
if idleTimeout != 0 {
idleTimeout += 5 * time.Second
}
return idleTimeout
}
/*
* 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 proxy
import (
"context"
"io"
"log/slog"
"net/http"
"reflect"
"strings"
"github.com/gravitational/trace"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/apis/meta/v1/unstructured"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/apimachinery/pkg/runtime/schema"
"k8s.io/client-go/dynamic"
"k8s.io/client-go/rest"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/utils/set"
"github.com/gravitational/teleport/lib/utils/slices"
)
// deleteResourcesCollection calls listResources method to list the resources the user
// has access to and calls their delete method using the allowed kube principals.
func (f *Forwarder) deleteResourcesCollection(sess *clusterSession, w http.ResponseWriter, req *http.Request) (resp any, err error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/deleteResourcesCollection",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("deleteResourcesCollection"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
req = req.WithContext(ctx)
// status holds the returned response code.
var status int
switch {
// Check if the target Kubernetes cluster is not served by the current service.
// If it's the case, forward the request to the target Kube Service where the
// filtering logic will be applied.
case !sess.isLocalKubernetesCluster:
rw := httplib.NewResponseStatusRecorder(w)
sess.forwarder.ServeHTTP(rw, req)
status = rw.Status()
default:
memoryRW := responsewriters.NewMemoryResponseWriter()
listReq := req.Clone(req.Context())
// reset body and method since list does not need the body response.
listReq.Body = nil
listReq.Method = http.MethodGet
_, err = f.listResources(sess, memoryRW, listReq)
if err != nil {
return nil, trace.Wrap(err)
}
// decompress the response body to be able to parse it.
if err := decompressInplace(memoryRW); err != nil {
return nil, trace.Wrap(err)
}
status, err = f.handleDeleteCollectionReq(req, sess, memoryRW, w)
if err != nil {
return nil, trace.Wrap(err)
}
}
f.emitAuditEvent(req, sess, status)
return nil, nil
}
func (f *Forwarder) handleDeleteCollectionReq(req *http.Request, sess *clusterSession, memWriter *responsewriters.MemoryResponseWriter, w http.ResponseWriter) (int, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/handleDeleteCollectionReq",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("deletePodsCollection"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
const internalErrStatus = http.StatusInternalServerError
// get content-type value
deleteRequestContentType := responsewriters.GetContentTypeHeader(req.Header)
deleteRequestEncoder, deleteRequestDecoder, err := newEncoderAndDecoderForContentType(
deleteRequestContentType,
newClientNegotiator(sess.codecFactory),
)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
deleteOptions, err := parseDeleteCollectionBody(req.Body, deleteRequestDecoder)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
req.Body.Close()
// decode memory rw body.
// We are reading an API request and API honors the GVK in the request so we don't
// need to set it.
_, listRequestDecoder, err := newEncoderAndDecoderForContentType(
responsewriters.GetContentTypeHeader(memWriter.Header()),
newClientNegotiator(sess.codecFactory),
)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
obj, err := decodeAndSetGVK(listRequestDecoder, memWriter.Buffer().Bytes(), nil /* defaults GVK */)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
details, err := f.findKubeDetailsByClusterName(sess.kubeClusterName)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
params := deleteResourcesCommonParams{
ctx: ctx,
log: f.log,
authCtx: &sess.authContext,
header: req.Header,
kubeDetails: details,
}
// At this point, items already include the list of pods the filtered pods the
// user has access to.
// For each Pod, we compute the kubernetes_groups and kubernetes_labels
// that are applicable and we will forward them as the delete request.
// If request is a dry-run.
// TODO (tigrato):
// - parallelize loop
// - check if the request should stop at the first fail.
switch o := obj.(type) {
case *metav1.Status:
// Do nothing.
case *unstructured.Unstructured:
if !o.IsList() {
return internalErrStatus, trace.BadParameter("unexpected CRD type")
}
list, err := o.ToList()
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
items, err := deleteResources(
params,
sess.metaResource.requestedResource.resourceKind,
sess.metaResource.requestedResource.apiGroup,
o.GetAPIVersion(),
slices.ToPointers(list.Items),
deleteOptions,
)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
objList := make([]any, 0, len(items))
for _, item := range items {
objList = append(objList, item.Object)
}
o.Object["items"] = objList
default:
output, err := getItemsUsingReflection(obj)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
if len(output.items) == 0 {
break
}
apiVersion, itemsR, objs, underlyingType := output.apiVersion, output.underlyingValue, output.items, output.underlyingType
items, err := deleteResources(
params,
sess.metaResource.requestedResource.resourceKind,
sess.metaResource.requestedResource.apiGroup,
apiVersion,
objs,
deleteOptions,
)
if err != nil {
return internalErrStatus, trace.Wrap(err)
}
setItemsUsingReflection(itemsR, underlyingType, items)
}
// reset the memory buffer.
memWriter.Buffer().Reset()
// set the content type to the response writer to match the delete
// request content type instead of the list request content type.
memWriter.Header().Set(
responsewriters.ContentTypeHeader,
deleteRequestContentType,
)
// encode the filtered response into the memory buffer.
if err := deleteRequestEncoder.Encode(obj, memWriter.Buffer()); err != nil {
return internalErrStatus, trace.Wrap(err)
}
// copy the output into the user's ResponseWriter and return.
return memWriter.Status(), trace.Wrap(memWriter.CopyInto(w))
}
type getItemsUsingReflectionOutput struct {
items []kubeObjectInterface
apiVersion string
underlyingType reflect.Type
underlyingValue reflect.Value
}
func getItemsUsingReflection(obj runtime.Object) (getItemsUsingReflectionOutput, error) {
// itemsFieldName is the field name of the items in the list
// object. This is used to get the items from the list object.
// We use reflection to get the items field name since
// the list object can be of any type.
const itemsFieldName = "Items"
objReflect := reflect.ValueOf(obj).Elem()
itemsR := objReflect.FieldByName(itemsFieldName)
if itemsR.Type().Kind() != reflect.Slice {
return getItemsUsingReflectionOutput{}, trace.BadParameter("unexpected type %T, Items is not a slice", obj)
}
if itemsR.Len() == 0 {
return getItemsUsingReflectionOutput{}, nil
}
var (
underlyingType = itemsR.Index(0).Type()
apiVersion, _ = obj.GetObjectKind().GroupVersionKind().ToAPIVersionAndKind()
objs = make([]kubeObjectInterface, 0, itemsR.Len())
)
for i := range itemsR.Len() {
item := itemsR.Index(i).Addr().Interface()
if item, ok := item.(kubeObjectInterface); ok {
objs = append(objs, item)
} else {
return getItemsUsingReflectionOutput{}, trace.BadParameter("unexpected type %T", itemsR.Interface())
}
}
return getItemsUsingReflectionOutput{
items: objs,
apiVersion: apiVersion,
underlyingType: underlyingType,
underlyingValue: itemsR,
}, nil
}
func setItemsUsingReflection(itemsR reflect.Value, underlyingType reflect.Type, items []kubeObjectInterface) {
// make a new slice of the same type as the original one.
slice := reflect.MakeSlice(itemsR.Type(), len(items), len(items))
for i, item := range items {
// convert the item to the underlying type of the slice.
// this is needed because items is a slice of pointers that
// satisfy the kubeObjectInterface interface.
// but the underlying type of the slice of elements is not
// a pointer. We dereference the item and convert it to the
// original slice element type.
slice.Index(i).Set(reflect.ValueOf(item).Elem().Convert(underlyingType))
}
itemsR.Set(slice)
}
// newImpersonatedKubeClient creates a new Kubernetes Client that impersonates
// a username and the groups.
func newImpersonatedKubeClient(creds kubeCreds, username string, groups []string) (*dynamic.DynamicClient, error) {
// clone cluster's rest config.
c := *creds.getKubeRestConfig()
// change the impersonated headers.
c.Impersonate = rest.ImpersonationConfig{
UserName: username,
Groups: groups,
}
client, err := dynamic.NewForConfig(&c)
return client, trace.Wrap(err)
}
// parseDeleteCollectionBody parses the request body targeted to pod collection
// endpoints.
func parseDeleteCollectionBody(r io.Reader, decoder runtime.Decoder) (metav1.DeleteOptions, error) {
into := metav1.DeleteOptions{}
data, err := io.ReadAll(r)
if err != nil {
return into, trace.Wrap(err)
}
if len(data) == 0 {
return into, nil
}
_, _, err = decoder.Decode(data, nil, &into)
return into, trace.Wrap(err)
}
type deleteResourcesCommonParams struct {
ctx context.Context
log *slog.Logger
authCtx *authContext
header http.Header
kubeDetails *kubeDetails
}
func deleteResources[T kubeObjectInterface](
params deleteResourcesCommonParams,
kind, group, apiVersion string,
items []T,
deleteOptions metav1.DeleteOptions,
) ([]T, error) {
deletedItems := make([]T, 0, len(items))
checker, err := params.authCtx.getAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
for _, item := range items {
// Compute users and groups from available roles that match the
// cluster labels and kubernetes resources.
allowedKubeGroups, allowedKubeUsers, err := checker.Kube().GetGroupsAndUsers(
params.authCtx.sessionTTL,
false,
services.NewKubernetesClusterLabelMatcher(
params.authCtx.kubeClusterLabels,
checker.AccessInfo().Username,
params.authCtx.CheckerContext.Traits(),
),
services.NewKubernetesResourceMatcher(
getKubeResource(kind, group, types.KubeVerbDeleteCollection, item),
params.authCtx.metaResource.isClusterWideResource(),
),
)
// no match was found, we ignore the request.
if err != nil {
continue
}
allowedKubeUsers, allowedKubeGroups = fillDefaultKubePrincipalDetails(allowedKubeUsers, allowedKubeGroups, params.authCtx.User.GetName())
impersonatedUsers, impersonatedGroups, err := computeAndValidateImpersonatedPrincipals(
set.New(allowedKubeUsers...), set.New(allowedKubeGroups...),
params.authCtx.User.GetName(),
params.header,
)
if err != nil {
continue
}
// create a new kubernetes.Client using the impersonated users and groups
// that matched the current pod.
client, err := newImpersonatedKubeClient(params.kubeDetails.kubeCreds, impersonatedUsers, impersonatedGroups)
if err != nil {
return nil, trace.Wrap(err)
}
gvk := item.GroupVersionKind()
if gvk.Group == "" || gvk.Version == "" {
tmp := strings.Split(apiVersion, "/")
if len(tmp) == 2 {
gvk.Group = tmp[0]
gvk.Version = tmp[1]
} else {
gvk.Version = apiVersion
}
}
// delete each resource individually.
err = client.Resource(schema.GroupVersionResource{
Group: gvk.Group,
Version: gvk.Version,
Resource: kind,
}).Namespace(item.GetNamespace()).Delete(params.ctx, item.GetName(), deleteOptions)
if err != nil {
// TODO(tigrato): check what should we do when delete returns an error.
// Should we check if it's permission error?
// Check if the Pod has already been deleted by a concurrent request
continue
}
deletedItems = append(deletedItems, item)
}
return deletedItems, nil
}
/*
* 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 proxy
import (
"bytes"
"context"
"io"
"log/slog"
"mime"
"net/http"
"github.com/gravitational/trace"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/apis/meta/v1/unstructured"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/apimachinery/pkg/runtime/schema"
"k8s.io/apimachinery/pkg/runtime/serializer"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
)
// needsFiltering returns true if RBAC filtering is required for the given rules.
func needsFiltering(allowedResources, deniedResources []types.KubernetesResource) bool {
return !containsWildcard(allowedResources) || len(deniedResources) != 0
}
// newResourceFilterer creates a wrapper function that once executed creates
// a runtime filter for kubernetes resources.
// The filter exclusion criteria is:
// - deniedResources: excluded if (namespace,name) matches an entry even if it matches
// the allowedResources's list.
// - allowedResources: excluded if (namespace,name) not match a single entry.
func newResourceFilterer(mr metaResource, codecs *serializer.CodecFactory, matcher resourceMatcher, log *slog.Logger) responsewriters.FilterWrapper {
return func(contentType string, responseCode int) (responsewriters.Filter, error) {
negotiator := newClientNegotiator(codecs)
encoder, decoder, err := newEncoderAndDecoderForContentType(contentType, negotiator)
if err != nil {
return nil, trace.Wrap(err)
}
return &resourceFilterer{
encoder: encoder,
decoder: decoder,
contentType: contentType,
responseCode: responseCode,
negotiator: negotiator,
log: log,
metaResource: mr,
matcher: matcher,
}, nil
}
}
// wildcardFilter is a filter that matches all pods.
var wildcardFilter = types.KubernetesResource{
Kind: types.Wildcard,
APIGroup: types.Wildcard,
Namespace: types.Wildcard,
Name: types.Wildcard,
Verbs: []string{types.Wildcard},
}
// containsWildcard returns true if the list of resources contains a wildcard filter.
func containsWildcard(resources []types.KubernetesResource) bool {
for _, r := range resources {
if r.Kind == wildcardFilter.Kind &&
r.APIGroup == wildcardFilter.APIGroup &&
r.Name == wildcardFilter.Name &&
r.Namespace == wildcardFilter.Namespace &&
len(r.Verbs) == 1 && r.Verbs[0] == wildcardFilter.Verbs[0] {
return true
}
}
return false
}
// resourceFilterer is a resource filterer instance.
type resourceFilterer struct {
encoder runtime.Encoder
decoder runtime.Decoder
// contentType is the response "Content-Type" header.
contentType string
// responseCode is the response status code.
responseCode int
// negotiator is an instance of a client negotiator.
negotiator runtime.ClientNegotiator
// log is the logger.
log *slog.Logger
// metaResource contains the information about the resource being filtered.
metaResource metaResource
// matcher is the per-item RBAC matcher (either fast precompiled or fallback per-item).
matcher resourceMatcher
}
// resourceMatcher matches a Kubernetes resource by name and namespace.
//
// A Teleport RBAC rule has five fields: kind, verb, apiGroup, namespace, and name.
// Only name and namespace vary per item in a list response.
// The rest are constant for the entire request (determined by the URL and HTTP method),
// so can be resolved once when the matcher is constructed.
type resourceMatcher interface {
Match(name, namespace string) (bool, error)
}
func newMatcher(mr metaResource, allowed, denied []types.KubernetesResource, log *slog.Logger) resourceMatcher {
// The fast matcher cannot handle namespace special cases in KubeResourceMatchesRegex
// (read-only namespace visibility, namespace kind matching with different target selection).
if mr.requestedResource.resourceKind != "namespaces" {
fm, err := newFastMatcher(mr, allowed, denied)
if err != nil {
log.DebugContext(context.Background(), "Failed to compile fast matcher, falling back to per-item matching", "error", err)
} else {
return fm
}
}
return &defaultMatcher{
kind: mr.requestedResource.resourceKind,
verb: mr.verb,
apiGroup: mr.requestedResource.apiGroup,
isClusterWide: mr.isClusterWideResource(),
allowedResources: allowed,
deniedResources: denied,
}
}
// FilterBuffer receives a byte array, decodes the response into the appropriate
// type and filters the resources based on allowed and denied rules configured.
// After filtering them, it serializes the response and dumps it into output buffer.
// If any error occurs, the call returns an error.
func (d *resourceFilterer) FilterBuffer(buf []byte, output io.Writer) error {
// decode the response into the appropriate Kubernetes API type.
obj, bf, err := d.decode(buf)
if err != nil {
return trace.Wrap(err)
}
// if bf is not empty, it means that response does not contain any valid response
// and it should be safe to write it back into the buffer.
if len(bf) > 0 {
_, err = output.Write(buf)
return trace.Wrap(err)
}
if allowed, isList, err := d.FilterObj(obj); err != nil {
return trace.Wrap(err)
} else if !isList && !allowed {
// if the object is not a list and it's not allowed, then we should
// return an error.
return trace.AccessDenied("access denied")
}
// encode the filterer response back to the user.
return d.encode(obj, output)
}
// FilterObj receives a runtime.Object type and filters the resources on it
// based on allowed and denied rules.
// After filtering them, the obj is manipulated to hold the filtered information.
// The isAllowed boolean returned indicates if the client is allowed to receive the event
// with the object.
// The isListObj boolean returned indicates if the object is a list of resources.
func (d *resourceFilterer) FilterObj(obj runtime.Object) (isAllowed bool, isList bool, err error) {
ctx := context.Background()
switch o := obj.(type) {
case *metav1.Status:
// Status object is returned when the Kubernetes API returns an error and
// should be forwarded to the user.
return true, false, nil
case *unstructured.Unstructured:
if o.IsList() {
hasElemts := d.filterUnstructuredList(o)
return hasElemts, true, nil
}
result, err := d.matcher.Match(o.GetName(), o.GetNamespace())
if err != nil {
d.log.WarnContext(ctx, "Unable to compile regex expressions within kubernetes_resources", "error", err)
}
// if err is not nil or result is false, we should not include it.
return result, false, nil
case *metav1.Table:
_, err := d.filterMetaV1Table(o)
if err != nil {
return false, false, trace.Wrap(err)
}
return len(o.Rows) > 0, true, nil
default:
if _, ok := obj.(metav1.ListInterface); ok {
output, err := getItemsUsingReflection(obj)
if err != nil {
return false, false, trace.Wrap(err, "failed to get items from list object")
}
if len(output.items) > 0 {
output.items = filterResourceList(d, output.items)
setItemsUsingReflection(output.underlyingValue, output.underlyingType, output.items)
}
return len(output.items) > 0, true, nil
} else if kubeObj, ok := o.(kubeObjectInterface); ok {
result, err := d.filterResource(kubeObj)
if err != nil {
d.log.WarnContext(ctx, "Unable to compile regex expressions within kubernetes_resources", "error", err)
}
// if err is not nil or result is false, we should not include it.
return result, false, nil
}
// It's important default types are never blindly forwarded or protocol
// extensions could result in information disclosures.
return false, false, trace.BadParameter("unexpected type received; got %T", obj)
}
}
// decode decodes the buffer into the appropriate type if the responseCode
// belongs to the range 200(OK)-206(PartialContent).
// If it does not belong, it returns the buffer unchanged since it contains
// an error message from the Kubernetes API server and it's safe to return
// it back to the user.
func (d *resourceFilterer) decode(buffer []byte) (runtime.Object, []byte, error) {
switch {
case d.responseCode == http.StatusSwitchingProtocols:
// no-op, we've been upgraded
return nil, buffer, nil
case d.responseCode < http.StatusOK /* 200 */ || d.responseCode > http.StatusPartialContent /* 206 */ :
// calculate an unstructured error from the response which the Result object may use if the caller
// did not return a structured error.
// Logic from: https://github.com/kubernetes/client-go/blob/58ff029093df37cad9fa28778a37f11fa495d9cf/rest/request.go#L1040
return nil, buffer, nil
default:
// We are reading an API request and API honors the GVK in the request so we don't
// need to set it.
out, err := decodeAndSetGVK(d.decoder, buffer, nil /* defaults GVK */)
return out, nil, trace.Wrap(err)
}
}
// decodePartialObjectMetadata decodes the metav1.PartialObjectMetadata present
// in the metav1.TableRow entry. This information comes from server side and
// includes the resource name and namespace as a structured object.
func (d *resourceFilterer) decodePartialObjectMetadata(row *metav1.TableRow) error {
if row.Object.Object != nil {
return nil
}
var err error
// decode only if row.Object.Object was not decoded before.
// We are reading an API request and API honors the GVK in the request so we don't
// need to set it.
row.Object.Object, err = decodeAndSetGVK(d.decoder, row.Object.Raw, nil /* defaults GVK */)
return trace.Wrap(err)
}
// encode encodes the filtered object into the io.Writer using the same
// content-type.
func (d *resourceFilterer) encode(obj runtime.Object, w io.Writer) error {
return trace.Wrap(d.encoder.Encode(obj, w))
}
// filterResourceList excludes resources the user should not have access to.
func filterResourceList[T kubeObjectInterface](d *resourceFilterer, originalList []T) []T {
filteredList := make([]T, 0, len(originalList))
for _, resource := range originalList {
if result, err := d.filterResource(resource); err == nil && result {
filteredList = append(filteredList, resource)
} else if err != nil {
slog.WarnContext(context.Background(), "Unable to compile regex expressions within kubernetes_resources", "error", err)
}
}
return filteredList
}
// kubeObjectInterface is an interface that all Kubernetes objects must
// implement to be able to filter them. It is used to extract the kind of the
// object from the GroupVersionKind object, the namespace and the name.
type kubeObjectInterface interface {
GroupVersionKind() schema.GroupVersionKind
GetNamespace() string
GetName() string
}
// filterResource validates if the user should access the current resource.
func (d *resourceFilterer) filterResource(resource kubeObjectInterface) (bool, error) {
return d.matcher.Match(resource.GetName(), resource.GetNamespace())
}
func getKubeResource(kind, group, verb string, obj kubeObjectInterface) types.KubernetesResource {
return types.KubernetesResource{
Kind: kind,
Namespace: obj.GetNamespace(),
Name: obj.GetName(),
Verbs: []string{verb},
APIGroup: group,
}
}
// filterMetaV1Table filters the serverside printed table to exclude resources
// that the user must not have access to.
func (d *resourceFilterer) filterMetaV1Table(table *metav1.Table) (*metav1.Table, error) {
resources := make([]metav1.TableRow, 0, len(table.Rows))
for i := range table.Rows {
row := &(table.Rows[i])
if err := d.decodePartialObjectMetadata(row); err != nil {
return nil, trace.Wrap(err)
}
resource, err := getKubeResourcePartialMetadataObject(d.metaResource.requestedResource.resourceKind, d.metaResource.requestedResource.apiGroup, d.metaResource.verb, row.Object.Object)
if err != nil {
return nil, trace.Wrap(err)
}
if result, err := d.matcher.Match(resource.Name, resource.Namespace); err != nil {
d.log.WarnContext(context.Background(), "Unable to compile regex expression", "error", err)
} else if result {
resources = append(resources, *row)
}
}
table.Rows = resources
return table, nil
}
// getKubeResourcePartialMetadataObject checks if obj satisfies namespaceNamer or namer interfaces
// otherwise returns an error.
func getKubeResourcePartialMetadataObject(kind, group, verb string, obj runtime.Object) (types.KubernetesResource, error) {
type namer interface {
GetName() string
}
type namespaceNamer interface {
GetNamespace() string
namer
}
switch o := obj.(type) {
case namespaceNamer:
return types.KubernetesResource{
Namespace: o.GetNamespace(),
Name: o.GetName(),
Kind: kind,
Verbs: []string{verb},
APIGroup: group,
}, nil
case namer:
return types.KubernetesResource{
Name: o.GetName(),
Kind: kind,
Verbs: []string{verb},
APIGroup: group,
}, nil
default:
return types.KubernetesResource{}, trace.BadParameter("unexpected %T type", obj)
}
}
// newEncoderAndDecoderForContentType creates a new encoder and decoder instances
// for the given contentType.
// If the contentType is invalid or not supported this function returns an error.
// Supported content types:
// - "application/json"
// - "application/yaml"
// - "application/vnd.kubernetes.protobuf"
func newEncoderAndDecoderForContentType(contentType string, negotiator runtime.ClientNegotiator) (runtime.Encoder, runtime.Decoder, error) {
mediaType, params, err := mime.ParseMediaType(contentType)
if err != nil {
return nil, nil, trace.WrapWithMessage(err, "unable to parse %q header %q", responsewriters.ContentTypeHeader, contentType)
}
dec, err := negotiator.Decoder(mediaType, params)
if err != nil {
return nil, nil, trace.Wrap(err)
}
enc, err := negotiator.Encoder(mediaType, params)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return enc, dec, nil
}
// decodeAndSetGVK decodes the payload into the appropriate type using the decoder
// provider and sets the GVK if available.
// defaults is the fallback GVK used by the decoder if the payload doesn't set their
// own GVK.
func decodeAndSetGVK(decoder runtime.Decoder, payload []byte, defaults *schema.GroupVersionKind) (runtime.Object, error) {
obj, gvk, err := decoder.Decode(payload, defaults, nil)
if err != nil {
return nil, trace.Wrap(err)
}
if gvk != nil {
// objects from decode do not contain GroupVersionKind.
// We force it to be present for later encoding.
obj.GetObjectKind().SetGroupVersionKind(*gvk)
}
return obj, nil
}
// filterBuffer filters the response buffer before writing it into the original
// MemoryResponseWriter.
func filterBuffer(filterWrapper responsewriters.FilterWrapper, src *responsewriters.MemoryResponseWriter) error {
if filterWrapper == nil {
return nil
}
filter, err := filterWrapper(responsewriters.GetContentTypeHeader(src.Header()), src.Status())
if err != nil {
return trace.Wrap(err)
}
// copy body into another slice so we can manipulate it.
b := bytes.NewBuffer(make([]byte, 0, src.Buffer().Len()))
// get the compressor and decompressor for the response based on the content type.
compressor, decompressor, err := getResponseCompressorDecompressor(src.Header())
if err != nil {
return trace.Wrap(err)
}
// decompress the response body into b.
if err := decompressor(b, src.Buffer()); err != nil {
return trace.Wrap(err)
}
// filter.FilterBuffer encodes the filtered payload into src.Buffer, so we need to
// reset it to discard the old payload.
src.Buffer().Reset()
// creates a compressor that writes the filtered payload into src.Buffer.
comp := compressor(src.Buffer())
// Close is a no-op operation into src but it's required to put the gzip writer
// into the sync.Pool.
defer comp.Close()
return trace.Wrap(filter.FilterBuffer(b.Bytes(), comp))
}
// filterUnstructuredList filters the unstructured list object to exclude resources
// that the user must not have access to.
// The filtered list is re-assigned to `obj.Object["items"]`.
func (d *resourceFilterer) filterUnstructuredList(obj *unstructured.Unstructured) (hasElems bool) {
const itemsKey = "items"
if obj == nil || obj.Object == nil {
return false
}
objList, err := obj.ToList()
if err != nil {
// This should never happen, but if it does, we should log it.
slog.WarnContext(context.Background(), "Unable to convert unstructured object to list", "error", err)
return false
}
filteredList := make([]any, 0, len(objList.Items))
for _, resource := range objList.Items {
if result, err := d.matcher.Match(resource.GetName(), resource.GetNamespace()); err != nil {
slog.WarnContext(context.Background(), "Unable to compile regex expressions within kubernetes_resources", "error", err)
} else if result {
filteredList = append(filteredList, resource.Object)
}
}
obj.Object[itemsKey] = filteredList
return len(filteredList) > 0
}
// defaultMatcher uses the existing matchKubernetesResource path for per-item matching.
type defaultMatcher struct {
kind string
verb string
apiGroup string
isClusterWide bool
allowedResources []types.KubernetesResource
deniedResources []types.KubernetesResource
}
func (m *defaultMatcher) Match(name, namespace string) (bool, error) {
resource := types.KubernetesResource{
Kind: m.kind,
Namespace: namespace,
Name: name,
Verbs: []string{m.verb},
APIGroup: m.apiGroup,
}
return matchKubernetesResource(resource, m.isClusterWide, m.allowedResources, m.deniedResources)
}
/*
* 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 proxy
import (
"bytes"
"context"
"io"
"net/http"
"strings"
"sync"
"time"
"github.com/gravitational/trace"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/utils"
)
// listResources forwards the pod list request to the target server, captures
// all output and filters accordingly to user roles resource access rules.
func (f *Forwarder) listResources(sess *clusterSession, w http.ResponseWriter, req *http.Request) (resp any, err error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/listResources",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("listResources"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
req = req.WithContext(ctx)
isLocalKubeCluster := sess.isLocalKubernetesCluster
supportsType := false
if isLocalKubeCluster {
_, supportsType = sess.rbacSupportedResources.getTeleportResourceKindFromAPIResource(sess.metaResource.requestedResource)
}
// status holds the returned response code.
var status int
defer func() {
if err != nil {
return
}
if status == 0 {
// Preserve pre-streaming behavior: treat unset status as 200 OK.
status = http.StatusOK
}
f.emitAuditEvent(req, sess, status)
}()
// Check if the target Kubernetes cluster is not served by the current service.
// If it's the case, forward the request to the target Kube Service where the
// filtering logic will be applied.
if !isLocalKubeCluster || !supportsType {
rw := httplib.NewResponseStatusRecorder(w)
sess.forwarder.ServeHTTP(rw, req)
status = rw.Status()
} else {
checker, err := sess.authContext.getAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
allowedResources, deniedResources := checker.Kube().GetResources(sess.kubeCluster)
shouldBeAllowed, err := matchListRequestShouldBeAllowed(sess.metaResource, allowedResources, deniedResources)
if err != nil {
return nil, trace.Wrap(err)
}
if !shouldBeAllowed {
notFoundMessage := f.kubeResourceDeniedAccessMsg(
sess.User.GetName(),
sess.metaResource.verb,
sess.metaResource.requestedResource,
)
return nil, trace.AccessDenied("%s", notFoundMessage)
}
// Identify if the request is long-lived watch stream based on
// HTTP connection.
if !isKubeWatchRequest(req, sess.authContext.metaResource.requestedResource) {
// List resources.
status, err = f.listResourcesList(req, w, sess, allowedResources, deniedResources)
} else {
// Creates a watch stream to the upstream target and applies filtering
// for each new frame that is received to exclude resources the user doesn't
// have access to.
status, err = f.listResourcesWatcher(req, w, sess, allowedResources, deniedResources)
}
if err != nil {
return nil, trace.Wrap(err)
}
}
return nil, nil
}
// listResourcesList forwards the request into the target cluster and accumulates the
// response into the memory. Once the request finishes, the memory buffer
// data is parsed and resources the user does not have access to are excluded from
// the response. Finally, the filtered response is serialized and sent back to
// the user with the appropriate headers.
func (f *Forwarder) listResourcesList(req *http.Request, w http.ResponseWriter, sess *clusterSession, allowedResources, deniedResources []types.KubernetesResource) (int, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/listResourcesList",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
req = req.WithContext(ctx)
if _, ok := sess.rbacSupportedResources.getTeleportResourceKindFromAPIResource(sess.metaResource.requestedResource); !ok {
return http.StatusBadRequest, trace.BadParameter("unknown resource kind %q", sess.metaResource.requestedResource.resourceKind)
}
// Check if filtering is needed before buffering the entire response.
// If the user has wildcard access and no denied resources, we can skip
// buffering and directly forward the response for better performance.
if !needsFiltering(allowedResources, deniedResources) {
// No filtering needed - use direct forwarding with status recording only.
// This avoids buffering the entire response in memory and the subsequent
// deserialization/re-serialization overhead.
rw := httplib.NewResponseStatusRecorder(w)
sess.forwarder.ServeHTTP(rw, req)
return rw.Status(), nil
}
// Filtering is needed. Use a filtering response writer that inspects headers
// and routes the body to either the streaming or buffered filter path.
matcher := newMatcher(sess.metaResource, allowedResources, deniedResources, f.log)
filterWrapper := newResourceFilterer(sess.metaResource, sess.codecFactory, matcher, f.log)
fw := newFilteringResponseWriter(w, matcher, filterWrapper, f.log, ctx, f.cfg.tracer, f.cfg.KubeServiceType)
sess.forwarder.ServeHTTP(fw, req)
return fw.Finish()
}
// matchListRequestShouldBeAllowed assess whether the user is permitted to perform its request
// based on the defined kubernetes_resource rules. The aim is to catch cases when the user
// has no access and present then a more user-friendly error message instead of returning
// an empty list.
// This function is not responsible for enforcing access rules.
func matchListRequestShouldBeAllowed(mr metaResource, allowedResources, deniedResources []types.KubernetesResource) (bool, error) {
resource := mr.rbacResource()
if resource == nil {
// Cluster is offline.
return false, nil
}
result, err := utils.KubeResourceCouldMatchRules(*resource, mr.isClusterWideResource(), deniedResources, types.Deny)
if err != nil {
return false, trace.Wrap(err)
} else if result {
return false, nil
}
result, err = utils.KubeResourceCouldMatchRules(*resource, mr.isClusterWideResource(), allowedResources, types.Allow)
if err != nil {
return false, trace.Wrap(err)
}
return result, nil
}
// listResourcesWatcher handles a long lived connection to the upstream server where
// the Kubernetes API returns frames with events.
// This handler creates a WatcherResponseWriter that spins a new goroutine once
// the API server writes the status code and headers.
// The goroutine waits for new events written into the response body and
// decodes each event. Once decoded, we validate if the Pod name matches
// any Pod specified in `kubernetes_resources` and if included, the event is
// forwarded to the user's response writer.
// If it does not match, the watcher ignores the event and continues waiting
// for the next event.
func (f *Forwarder) listResourcesWatcher(req *http.Request, w http.ResponseWriter, sess *clusterSession, allowedResources, deniedResources []types.KubernetesResource) (int, error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/listResourcesWatcher",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
req = req.WithContext(ctx)
negotiator := newClientNegotiator(sess.codecFactory)
_, ok := sess.rbacSupportedResources.getTeleportResourceKindFromAPIResource(sess.metaResource.requestedResource)
if !ok {
return http.StatusBadRequest, trace.BadParameter("unknown resource kind %q", sess.metaResource.requestedResource.resourceKind)
}
var filter responsewriters.FilterWrapper
if needsFiltering(allowedResources, deniedResources) {
filter = newResourceFilterer(
sess.metaResource,
sess.codecFactory,
newMatcher(sess.metaResource, allowedResources, deniedResources, f.log),
f.log,
)
}
rw, err := responsewriters.NewWatcherResponseWriter(
w,
negotiator,
filter,
)
if err != nil {
return http.StatusInternalServerError, trace.Wrap(err)
}
// if this pod watch request is for a specific pod, watch for and
// push events that show ephemeral containers were started if there
// are any ephemeral containers waiting to be created for this pod
// by this user
var wg sync.WaitGroup
ctx, cancel := context.WithCancel(req.Context())
if podName := isRequestTargetedToPod(req, sess.metaResource.requestedResource); podName != "" && ok {
wg.Add(1)
go func() {
defer wg.Done()
f.sendEphemeralContainerEvents(ctx, rw, sess, podName)
}()
}
// Forwards the request to the target cluster.
sess.forwarder.ServeHTTP(rw, req)
// Wait for the fake event pushing goroutine to finish
cancel()
wg.Wait()
// Once the request terminates, close the watcher and waits for resources
// cleanup.
err = rw.Close()
return rw.Status(), trace.Wrap(err)
}
// sendEphemeralContainerEvents will poll the list of ephemeral containers
// each 5s from cache and see if they match the user and pod and namespace.
// If any match exists, it will push a fake event to the watcher stream to trick
// kubectl into creating the exec session.
func (f *Forwarder) sendEphemeralContainerEvents(ctx context.Context, rw *responsewriters.WatcherResponseWriter, sess *clusterSession, podName string) {
const backoff = 5 * time.Second
sentDebugContainers := map[string]struct{}{}
ticker := time.NewTicker(backoff)
defer ticker.Stop()
for {
wcs, err := f.getUserEphemeralContainersForPod(
ctx,
sess.User.GetName(),
sess.kubeClusterName,
sess.metaResource.requestedResource.namespace,
podName,
)
if err != nil {
f.log.WarnContext(ctx, "error getting user ephemeral containers", "error", err)
return
}
for _, wc := range wcs {
if _, ok := sentDebugContainers[wc.GetSpec().GetContainerName()]; ok {
continue
}
evt, err := f.getPatchedPodEvent(ctx, sess, wc)
if err != nil {
f.log.WarnContext(ctx, "error pushing pod event", "error", err)
continue
}
sentDebugContainers[wc.GetSpec().GetContainerName()] = struct{}{}
// push the event to the client
// this will lock until the event is pushed or the
// request context is done.
rw.PushVirtualEventToClient(ctx, evt)
}
// wait a bit before querying the cache again, or return
// if the request has finished
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
// decompressInplace decompresses the response into the same buffer it was
// written to.
// If the response is not compressed, it does nothing.
func decompressInplace(memoryRW *responsewriters.MemoryResponseWriter) error {
switch memoryRW.Header().Get(contentEncodingHeader) {
case contentEncodingGZIP:
_, decompressor, err := getResponseCompressorDecompressor(memoryRW.Header())
if err != nil {
return trace.Wrap(err)
}
newBuf := bytes.NewBuffer(nil)
_, err = io.Copy(newBuf, memoryRW.Buffer())
if err != nil {
return trace.Wrap(err)
}
memoryRW.Buffer().Reset()
err = decompressor(memoryRW.Buffer(), newBuf)
return trace.Wrap(err)
default:
return nil
}
}
// isRequestTargetedToPod checks if the request is
// possibly targeted to an ephemeral container. If it is, it returns the
// name of the pod that the container is in.
// This function is used to determine if a watch request is for a specific pod
// because although the watch request is for a specific pod, the endpoint
// is the same as the endpoint for the pod list request.
// A request targeted to an ephemeral container will follow this template:
// GET api/v1/namespaces/<namespace>/pods?fieldSelector=metadata.name%3D<pod_name>
func isRequestTargetedToPod(req *http.Request, kube apiResource) string {
const podsResource = "pods"
if kube.resourceKind != podsResource {
return ""
}
if kube.namespace == "" {
return ""
}
if kube.resourceName != "" {
return ""
}
q := req.URL.Query()
fieldSel, ok := q["fieldSelector"]
if !ok {
return ""
}
for _, val := range fieldSel {
if podName, ok := strings.CutPrefix(val, "metadata.name="); ok {
return podName
}
}
return ""
}
/*
* 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 proxy
import (
"bytes"
"fmt"
"io"
"net/http"
"slices"
"strconv"
"github.com/gravitational/trace"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib/reverseproxy"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/services"
)
// rewriteResponseForbidden rewrites the response body when the response includes
// a GKE Autopilot forbidden error caused by impersonating system:masters group.
// The response body is rewritten to include a more user friendly error message.
// All other responses are returned as is.
// Example of response body that is rewritten:
//
// Error from server (Forbidden): groups "system:masters" is forbidden:
// User "<user>" cannot impersonate resource "groups" in API group "" at the cluster
// scope: GKE Warden authz [denied by user-impersonation-limitation]: impersonating
// system identities are not allowed
//
// The rewritten response body will look like:
//
// Error from server (Forbidden): "GKE Autopilot denied the request because it impersonates the "system:masters" group.
// Your Teleport Roles [role1,role2] have given access to the "system:masters" group for the cluster "<cluster>".
// For additional information and resolution, please visit
// https://goteleport.com/docs/enroll-resources/kubernetes-access/troubleshooting/#unable-to-connect-to-gke-autopilot-clusters
func (f *Forwarder) rewriteResponseForbidden(s *clusterSession) func(r *http.Response) error {
return func(r *http.Response) error {
const (
// The string that is returned by the GKE Autopilot cluster when
// users try to impersonate system:masters group.
autopilotForbidden = "impersonating system identities are not allowed"
)
// If the response is not forbidden, we don't need to do anything.
// The response will be returned as is and written to the client.
if r.StatusCode != http.StatusForbidden || r.Body == nil {
return nil
}
// create a new buffer to read the response body into.
b := bytes.NewBuffer(make([]byte, 0, 4096))
// Read the response body into the buffer.
if _, err := io.Copy(b, r.Body); err != nil {
return trace.Wrap(err)
}
// Close the response body.
if err := r.Body.Close(); err != nil {
return trace.Wrap(err)
}
// Replace the response body with the new buffer.
r.Body = io.NopCloser(b)
switch {
case bytes.Contains(b.Bytes(), []byte(autopilotForbidden)):
// If the response body contains the forbidden string, we rewrite the
// response body to include a more user friendly error message.
encoder, _, err := newEncoderAndDecoderForContentType(
r.Header.Get(responsewriters.ContentTypeHeader),
newClientNegotiator(&globalKubeCodecs),
)
if err != nil {
f.log.ErrorContext(r.Request.Context(), "Failed to create encoder", "error", err)
return nil
}
status := &metav1.Status{
Status: metav1.StatusFailure,
Code: int32(http.StatusForbidden),
Reason: metav1.StatusReasonForbidden,
Message: "GKE Autopilot denied the request because it impersonates the \"system:masters\" group.\n" +
fmt.Sprintf(
"Your Teleport Roles %v have given access to the \"system:masters\" group "+
"for the cluster %q.\n", collectSystemMastersTeleportRoles(s), s.kubeClusterName) +
"For additional information and resolution, " +
"please visit https://goteleport.com/docs/enroll-resources/kubernetes-access/troubleshooting/#unable-to-connect-to-gke-autopilot-clusters\n",
}
// Reset the buffer to write the new response.
b.Reset()
// Encode the new response.
if err = encoder.Encode(status, b); err != nil {
f.log.ErrorContext(r.Request.Context(), "Failed to encode response", "error", err)
return trace.Wrap(err)
}
// This function rewrote the response body, so we need update delete the
// Content-Length header to avoid mismatch between the actual body
// length and the original Content-Length header value.
r.Header.Set(reverseproxy.ContentLength, strconv.Itoa(b.Len()))
return nil
}
return nil
}
}
// collectSystemMastersTeleportRoles returns a list of teleport roles that grant
// system:masters to the target cluster.
func collectSystemMastersTeleportRoles(s *clusterSession) []string {
const (
systemMastersGroup = "system:masters"
)
accessChecker, err := s.authContext.getAccessChecker()
if err != nil {
return nil
}
matchers := make([]services.RoleMatcher, 0, 3)
// Creates a matcher that matches the cluster labels against `kubernetes_labels`
// defined for each user's role.
matchers = append(
matchers,
services.NewKubernetesClusterLabelMatcher(s.kubeClusterLabels, accessChecker.AccessInfo().Username, accessChecker.Traits()),
)
// If the kubeResource is available, append an extra matcher that validates
// if the kubernetes resource is allowed by the user roles that satisfy the
// target cluster labels.
// Each role defines `kubernetes_resources` and when kubeResource is available,
// KubernetesResourceMatcher will match roles that statisfy the resources at the
// same time that ClusterLabelMatcher matches the role's "kubernetes_labels".
// The call to roles.CheckKubeGroupsAndUsers when both matchers are provided
// results in the intersection of roles that match the "kubernetes_labels" and
// roles that allow access to the desired "kubernetes_resource".
// If from the intersection results an empty set, the request is denied.
if rbacRes := s.metaResource.rbacResource(); rbacRes != nil && !s.metaResource.isList {
matchers = append(
matchers,
services.NewKubernetesResourceMatcher(*rbacRes, s.metaResource.isClusterWideResource()),
)
}
var rolesWithSystemMasters []string
matchers = append(matchers,
// Creates a matcher that checks if the role grants system:masters group.
// The matcher will be called for each role that matches the cluster labels
// and the kubernetes resource (if available).
// It's important to note that this matcher must be the last one in the list
// otherwise the returned roles may not match the cluster labels and the
// kubernetes resource.
services.RoleMatcherFunc(func(r types.Role, cond types.RoleConditionType) (bool, error) {
groups := r.GetKubeGroups(cond)
if slices.Contains(groups, systemMastersGroup) {
rolesWithSystemMasters = append(rolesWithSystemMasters, r.GetName())
}
return true, nil
}),
)
_, _, _ = accessChecker.Kube().GetGroupsAndUsers(s.sessionTTL, false /* overrideTTL */, matchers...)
return rolesWithSystemMasters
}
/*
Copyright 2015 The Kubernetes Authors.
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.
*/
package proxy
import (
"bufio"
"context"
"crypto/tls"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"strings"
"time"
"github.com/gravitational/trace"
apierrors "k8s.io/apimachinery/pkg/api/errors"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/apimachinery/pkg/runtime/serializer"
utilnet "k8s.io/apimachinery/pkg/util/net"
"k8s.io/apimachinery/third_party/forked/golang/netutil"
"k8s.io/streaming/pkg/httpstream"
streamspdy "k8s.io/streaming/pkg/httpstream/spdy"
apiclient "github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/lib/kube/internal"
)
// SpdyRoundTripper knows how to upgrade an HTTP request to one that supports
// multiplexed streams. After RoundTrip() is invoked, the upgraded connection
// can be obtained by passing the response to NewConnection. SpdyRoundTripper
// implements the UpgradeRoundTripper interface.
//
// A SpdyRoundTripper is single-use and not safe for concurrent use:
// callers create one, use it for a single RoundTrip, and discard it.
// Its conn is tied to that one request rather than kept in a per-request map.
type SpdyRoundTripper struct {
roundTripperConfig
// conn is the underlying network connection to the remote server.
conn net.Conn
// cleanups contains objects that should be closed when the roundtripper is
// no longer used. [SpdyRoundTripper.Cleanup] should be called to ensure
// that.
cleanups []io.Closer
}
// Cleanup ensures that every connection that was opened by this roundtripper is
// closed.
func (w *SpdyRoundTripper) Cleanup() {
for _, closer := range w.cleanups {
_ = closer.Close()
}
w.cleanups = nil
}
var (
_ utilnet.TLSClientConfigHolder = &SpdyRoundTripper{}
_ httpstream.UpgradeRoundTripper = &SpdyRoundTripper{}
_ utilnet.Dialer = &SpdyRoundTripper{}
)
type roundTripperConfig struct {
// ctx is a context for this round tripper
ctx context.Context
// sess is the cluster session
sess *clusterSession
// dialWithContext is the function used connect to remote address
dialWithContext dialContextFunc
// tlsConfig holds the TLS configuration settings to use when connecting
// to the remote server.
tlsConfig *tls.Config
// pingPeriod is the period at which to send pings to the remote server to
// keep the SPDY connection alive.
pingPeriod time.Duration
// originalHeaders are the headers that were passed from the original request.
// These headers are used to set the headers on the new request if the user
// requested Kubernetes impersonation.
originalHeaders http.Header
// useIdentityForwarding controls whether the proxy should forward the
// identity of the user making the request to the remote server using the
// auth.TeleportImpersonateUserHeader and auth.TeleportImpersonateIPHeader
// headers instead of relying on the certificate to transport it.
useIdentityForwarding bool
// log specifies the logger.
log *slog.Logger
proxier func(*http.Request) (*url.URL, error)
}
// NewSpdyRoundTripperWithDialer creates a new SpdyRoundTripper that will use
// the specified tlsConfig. This function is mostly meant for unit tests.
func NewSpdyRoundTripperWithDialer(cfg roundTripperConfig) *SpdyRoundTripper {
return &SpdyRoundTripper{roundTripperConfig: cfg}
}
// TLSClientConfig implements pkg/util/net.TLSClientConfigHolder for proper TLS checking during
// proxying with a spdy roundtripper.
func (s *SpdyRoundTripper) TLSClientConfig() *tls.Config {
return s.tlsConfig
}
// Dial implements k8s.io/apimachinery/pkg/util/net.Dialer.
func (s *SpdyRoundTripper) Dial(req *http.Request) (net.Conn, error) {
conn, err := s.dial(req)
if err != nil {
return nil, err
}
if err := req.Write(conn); err != nil {
conn.Close()
return nil, err
}
return conn, nil
}
// dial dials the host specified by url, using TLS if appropriate.
func (s *SpdyRoundTripper) dial(req *http.Request) (conn net.Conn, err error) {
var proxyURL *url.URL
if s.proxier != nil {
proxyURL, err = s.proxier(req)
if err != nil {
return nil, err
}
}
if proxyURL == nil {
conn, err = s.dialWithoutProxy(req.URL)
} else {
conn, err = s.dialWithProxy(req, proxyURL)
}
if err != nil {
return nil, trace.Wrap(err)
}
if req.URL.Scheme == "https" {
return s.tlsConn(s.ctx, conn, netutil.CanonicalAddr(req.URL))
}
return conn, nil
}
func (s *SpdyRoundTripper) dialWithoutProxy(url *url.URL) (conn net.Conn, err error) {
dialAddr := netutil.CanonicalAddr(url)
switch {
case s.dialWithContext != nil:
conn, err = s.dialWithContext(s.ctx, "tcp", dialAddr)
default:
conn, err = net.Dial("tcp", dialAddr)
}
return conn, trace.Wrap(err)
}
// tlsConn returns a TLS client side connection using rwc as the underlying transport.
func (s *SpdyRoundTripper) tlsConn(ctx context.Context, rwc net.Conn, targetHost string) (net.Conn, error) {
host, _, err := net.SplitHostPort(targetHost)
if err != nil {
return nil, err
}
tlsConfig := s.tlsConfig
switch {
case tlsConfig == nil:
tlsConfig = &tls.Config{ServerName: host}
case len(tlsConfig.ServerName) == 0:
tlsConfig = tlsConfig.Clone()
tlsConfig.ServerName = host
}
tlsConn := tls.Client(rwc, tlsConfig)
// Client handshake will verify the server hostname and cert chain. That
// way we can err our before first read/write.
if err := tlsConn.HandshakeContext(ctx); err != nil {
tlsConn.Close()
return nil, trace.Wrap(err)
}
return tlsConn, nil
}
// dialWithProxy dials the host specified by url through an http or an socks5 proxy.
func (s *SpdyRoundTripper) dialWithProxy(req *http.Request, proxyURL *url.URL) (net.Conn, error) {
// ensure we use a canonical host with proxyReq
targetHost := netutil.CanonicalAddr(req.URL)
proxyDialConn, err := apiclient.DialProxyWithDialer(
s.ctx,
proxyURL,
targetHost,
apiclient.ContextDialerFunc(s.dialWithContext),
)
return proxyDialConn, trace.Wrap(err)
}
// RoundTrip executes the Request and upgrades it. After a successful upgrade,
// clients may call SpdyRoundTripper.NewConnection() to retrieve the upgraded connection.
func (s *SpdyRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
header := utilnet.CloneHeader(req.Header)
// copyImpersonationHeaders copies the headers from the original request to the new
// request headers. This is necessary to forward the original user's impersonation
// when multiple kubernetes_users are available.
copyImpersonationHeaders(header, s.originalHeaders)
header.Set(httpstream.HeaderConnection, httpstream.HeaderUpgrade)
header.Set(httpstream.HeaderUpgrade, streamspdy.HeaderSpdy31)
if err := setupImpersonationHeaders(s.sess, header); err != nil {
return nil, trace.Wrap(err)
}
// If we're using identity forwarding, we need to add the impersonation
// headers to the request before we send the request.
if s.useIdentityForwarding {
h, err := internal.IdentityForwardingHeaders(s.ctx, header)
if err != nil {
return nil, trace.Wrap(err)
}
header = h
}
clone := utilnet.CloneRequest(req)
clone.Header = header
conn, err := s.Dial(clone)
if err != nil {
return nil, err
}
responseReader := bufio.NewReader(conn)
resp, err := http.ReadResponse(responseReader, nil)
if err != nil {
conn.Close()
return nil, err
}
s.cleanups = append(s.cleanups, conn)
s.conn = conn
return resp, nil
}
// NewConnection validates the upgrade response, creating and returning a new
// httpstream.Connection if there were no errors.
func (s *SpdyRoundTripper) NewConnection(resp *http.Response) (httpstream.Connection, error) {
if s.conn == nil {
return nil, trace.Wrap(&upgradeFailureError{
Cause: errors.New("unable to upgrade connection: broken roundtripper setup, connection is missing but it should be present (this is a bug)"),
})
}
connectionHeader := strings.ToLower(resp.Header.Get(httpstream.HeaderConnection))
upgradeHeader := strings.ToLower(resp.Header.Get(httpstream.HeaderUpgrade))
if (resp.StatusCode != http.StatusSwitchingProtocols) ||
!strings.Contains(connectionHeader, strings.ToLower(httpstream.HeaderUpgrade)) ||
!strings.Contains(upgradeHeader, strings.ToLower(streamspdy.HeaderSpdy31)) {
// The upgrade was rejected. Close the conn after reading the response
// error so that we don't leak an io.Copy goroutine.
defer s.conn.Close()
return nil, trace.Wrap(extractKubeAPIStatusFromReq(resp))
}
return streamspdy.NewClientConnectionWithPings(s.conn, s.pingPeriod)
}
// statusScheme is a minimal scheme registering only metav1.Status,
// used to decode the status returned in an upgrade error response (see extractKubeAPIStatusFromReq).
var statusScheme = runtime.NewScheme()
// ParameterCodec knows about query parameters used with the meta v1 API spec.
var statusCodecs = serializer.NewCodecFactory(statusScheme)
func init() {
statusScheme.AddUnversionedTypes(metav1.SchemeGroupVersion,
&metav1.Status{},
)
}
// extractKubeAPIStatusFromReq extracts the status from the response body and returns it as an error.
func extractKubeAPIStatusFromReq(rsp *http.Response) error {
defer func() {
_ = rsp.Body.Close()
}()
responseError := ""
responseErrorBytes, err := io.ReadAll(rsp.Body)
if err != nil {
responseError = "unable to read error from server response"
} else {
if obj, _, err := statusCodecs.UniversalDecoder().Decode(responseErrorBytes, nil, &metav1.Status{}); err == nil {
if status, ok := obj.(*metav1.Status); ok {
return &upgradeFailureError{Cause: &apierrors.StatusError{ErrStatus: *status}}
}
}
responseError = string(responseErrorBytes)
responseError = strings.TrimSpace(responseError)
}
return &upgradeFailureError{Cause: fmt.Errorf("unable to upgrade connection: %s", responseError)}
}
// upgradeFailureError encapsulates the cause for why the streaming
// upgrade request failed. Implements error interface.
type upgradeFailureError struct {
Cause error
}
func (u *upgradeFailureError) Error() string {
return u.Cause.Error()
}
func (u *upgradeFailureError) Unwrap() error {
return u.Cause
}
func isTeleportUpgradeFailure(err error) bool {
var upgradeErr *upgradeFailureError
return errors.As(err, &upgradeErr)
}
/*
* Teleport
* Copyright (C) 2024 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 proxy
import (
"fmt"
"io"
"net/http"
gwebsocket "github.com/gorilla/websocket"
"github.com/gravitational/trace"
utilnet "k8s.io/apimachinery/pkg/util/net"
kwebsocket "k8s.io/client-go/transport/websocket"
"k8s.io/streaming/pkg/httpstream"
"github.com/gravitational/teleport/lib/kube/internal"
)
// WebsocketRoundTripper knows how to upgrade an HTTP request to one that supports
// multiplexed streams. After RoundTrip() is invoked, Conn will be set
// and usable. WebsocketRoundTripper implements the UpgradeRoundTripper interface.
type WebsocketRoundTripper struct {
roundTripperConfig
// conn is the websocket network connection to the remote server.
conn *gwebsocket.Conn
// cleanups contains objects that should be closed when the roundtripper is
// no longer used. Cleared by [WebsocketRoundTripper.Cleanup].
cleanups []io.Closer
// onConnected is a hook that happens when connection was successfully established,
// can be used to propagate established connection somewhere else - we are using it
// to set underlying connection of the native k8s websocket executor.
onConnected func(conn *gwebsocket.Conn)
}
// Cleanup ensures that every connection that was opened by this roundtripper is
// closed.
func (w *WebsocketRoundTripper) Cleanup() {
for _, closer := range w.cleanups {
_ = closer.Close()
}
w.cleanups = nil
}
// NewWebsocketRoundTripperWithDialer creates a new WebsocketRoundTripper that will
// dial and upgrade connection, copying impersonation setup specified in the config.
func NewWebsocketRoundTripperWithDialer(cfg roundTripperConfig) *WebsocketRoundTripper {
return &WebsocketRoundTripper{roundTripperConfig: cfg}
}
// RoundTrip executes the Request and upgrades it.
func (w *WebsocketRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) {
header := utilnet.CloneHeader(req.Header)
// copyImpersonationHeaders copies the headers from the original request to the new
// request headers. This is necessary to forward the original user's impersonation
// when multiple kubernetes_users are available.
copyImpersonationHeaders(header, w.originalHeaders)
if err := setupImpersonationHeaders(w.sess, header); err != nil {
return nil, trace.Wrap(err)
}
var err error
// If we're using identity forwarding, we need to add the impersonation
// headers to the request before we send the request.
if w.useIdentityForwarding {
if header, err = internal.IdentityForwardingHeaders(w.ctx, header); err != nil {
return nil, trace.Wrap(err)
}
}
clone := utilnet.CloneRequest(req)
clone.Header = header
nativeBufferSize := (&kwebsocket.RoundTripper{}).DataBufferSize()
wsDialer := gwebsocket.Dialer{
NetDialContext: w.dialWithContext,
Proxy: w.proxier,
TLSClientConfig: w.tlsConfig,
Subprotocols: header[httpstream.HeaderProtocolVersion],
ReadBufferSize: nativeBufferSize + 1024, // matching code in k8s websocket/roundripper.go
WriteBufferSize: nativeBufferSize + 1024,
}
switch clone.URL.Scheme {
case "https":
clone.URL.Scheme = "wss"
case "http":
clone.URL.Scheme = "ws"
default:
return nil, fmt.Errorf("unknown url scheme: %s", clone.URL.Scheme)
}
wsConn, wsResp, err := wsDialer.DialContext(w.ctx, clone.URL.String(), clone.Header)
if err != nil {
if wsResp != nil {
return nil, trace.Wrap(extractKubeAPIStatusFromReq(wsResp))
}
return nil, &httpstream.UpgradeFailureError{Cause: err}
}
w.cleanups = append(w.cleanups, wsConn)
w.conn = wsConn
if w.onConnected != nil {
w.onConnected(wsConn)
}
return wsResp, nil
}
/*
* 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 proxy
import (
"context"
"errors"
"log/slog"
"maps"
"slices"
"strings"
"github.com/gravitational/trace"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/apis/meta/v1/unstructured"
metav1beta1 "k8s.io/apimachinery/pkg/apis/meta/v1beta1"
"k8s.io/apimachinery/pkg/runtime"
"k8s.io/apimachinery/pkg/runtime/schema"
"k8s.io/apimachinery/pkg/runtime/serializer"
utilruntime "k8s.io/apimachinery/pkg/util/runtime"
"k8s.io/client-go/discovery"
"k8s.io/client-go/kubernetes"
"k8s.io/client-go/kubernetes/scheme"
"k8s.io/metrics/pkg/apis/metrics"
metricsv1beta1 "k8s.io/metrics/pkg/apis/metrics/v1beta1"
)
const (
// listSuffix is the suffix added to the name of the type to create the name
// of the list type.
// For example: "Role" -> "RoleList"
listSuffix = "List"
)
var (
// globalKubeScheme is the runtime Scheme that holds information about supported
// message types.
globalKubeScheme = runtime.NewScheme()
// globalKubeCodecs creates a serializer/deserizalier for the different codecs
// supported by the Kubernetes API.
globalKubeCodecs = serializer.NewCodecFactory(globalKubeScheme)
)
// Register all groups in the schema's registry.
// It manually registers support for `metav1.Table` because go-client does not
// support it but `kubectl` calls require support for it.
func init() {
// Register external types for Scheme
utilruntime.Must(registerDefaultKubeTypes(globalKubeScheme))
}
// registerDefaultKubeTypes registers the default types for the Kubernetes API into
// the given scheme.
func registerDefaultKubeTypes(s *runtime.Scheme) error {
// Register external types for Scheme
metav1.AddToGroupVersion(s, schema.GroupVersion{Group: "", Version: "v1"})
if err := metrics.AddToScheme(s); err != nil {
return trace.Wrap(err)
}
if err := metricsv1beta1.AddToScheme(s); err != nil {
return trace.Wrap(err)
}
if err := metav1.AddMetaToScheme(s); err != nil {
return trace.Wrap(err)
}
if err := metav1beta1.AddMetaToScheme(s); err != nil {
return trace.Wrap(err)
}
if err := scheme.AddToScheme(s); err != nil {
return trace.Wrap(err)
}
err := s.SetVersionPriority(corev1.SchemeGroupVersion)
return trace.Wrap(err)
}
// newClientNegotiator creates a negotiator that based on `Content-Type` header
// from the Kubernetes API response is able to create a different encoder/decoder.
// Supported content types:
// - "application/json"
// - "application/yaml"
// - "application/vnd.kubernetes.protobuf"
func newClientNegotiator(codecFactory *serializer.CodecFactory) runtime.ClientNegotiator {
return runtime.NewClientNegotiator(
codecFactory.WithoutConversion(),
schema.GroupVersion{
// create a serializer for Kube API v1
Version: "v1",
Group: "",
},
)
}
// gvkSupportedResourcesKey is the key used in gvkSupportedResources
// to map from a parsed API path to the corresponding resource GVK.
type gvkSupportedResourcesKey struct {
name string
apiGroup string
version string
}
// gvkSupportedResources maps a parsed API path to the corresponding resource GVK.
type gvkSupportedResources map[gvkSupportedResourcesKey]*schema.GroupVersionKind
// newClusterSchemaBuilder creates a new schema builder for the given cluster.
// This schema includes all well-known Kubernetes types and all namespaced
// custom resources.
// It also returns a map of resources that we support RBAC restrictions for.
func newClusterSchemaBuilder(log *slog.Logger, client kubernetes.Interface) (*serializer.CodecFactory, rbacSupportedResources, gvkSupportedResources, error) {
kubeScheme := runtime.NewScheme()
kubeCodecs := serializer.NewCodecFactory(kubeScheme)
supportedResources := make(rbacSupportedResources)
gvkSupportedRes := make(gvkSupportedResources)
if err := registerDefaultKubeTypes(kubeScheme); err != nil {
return nil, nil, nil, trace.Wrap(err)
}
// discoveryErr is returned when the discovery of one or more API groups fails.
var discoveryErr *discovery.ErrGroupDiscoveryFailed
// register all namespaced custom resources
_, apiGroups, err := client.Discovery().ServerGroupsAndResources()
switch {
case errors.As(err, &discoveryErr):
// If the discovery of one or more API groups fails, we still want to
// register the well-known Kubernetes types.
// This is because the discovery of API groups can fail if the APIService
// is not available. Usually, this happens when the API service is not local
// to the cluster (e.g. when API is served by a pod) and the service is not
// reachable.
// In this case, we still want to register the other resources that are
// available in the cluster.
log.DebugContext(context.Background(), "Failed to discover some API groups",
"groups", slices.Collect(maps.Keys(discoveryErr.Groups)),
"error", err,
)
case err != nil:
return nil, nil, nil, trace.Wrap(err)
}
for _, apiGroup := range apiGroups {
group, version := getKubeAPIGroupAndVersion(apiGroup.GroupVersion)
for _, apiResource := range apiGroup.APIResources {
// register all types
gvkSupportedRes[gvkSupportedResourcesKey{
name: apiResource.Name, /* pods, configmaps, ... */
apiGroup: group,
version: version,
}] = &schema.GroupVersionKind{
Group: group,
Version: version,
Kind: apiResource.Kind, /* Pod, ConfigMap ...*/
}
}
groupVersion := schema.GroupVersion{Group: group, Version: version}
for _, apiResource := range apiGroup.APIResources {
// build the resource key to be able to look it up later and check if
// if we support RBAC restrictions for it.
resourceKey := allowedResourcesKey{
apiGroup: group,
resourceKind: apiResource.Name,
}
supportedResources[resourceKey] = apiResource
// Create the group version kind for the resource.
gvk := groupVersion.WithKind(apiResource.Kind)
// Check if the resource is already registered in the scheme,
// if it is, we don't need to register it again.
if _, err := kubeScheme.New(gvk); err == nil {
continue
}
// Register the resource with the scheme to be able to decode it
// into an unstructured object.
kubeScheme.AddKnownTypeWithName(
gvk,
&unstructured.Unstructured{},
)
// Register the resource list with the scheme to be able to decode it
// into an unstructured object.
// Resource lists follow the naming convention: <resource-kind>List
kubeScheme.AddKnownTypeWithName(
groupVersion.WithKind(apiResource.Kind+listSuffix),
&unstructured.Unstructured{},
)
}
}
return &kubeCodecs, supportedResources, gvkSupportedRes, nil
}
// buildCodecsForGVKs builds a codec factory whose scheme knows the well-known
// Kubernetes types plus the given discovered GVKs (as unstructured). It mirrors
// newClusterSchemaBuilder's registration and is used to rebuild the codecs after
// a targeted, single-group-version discovery, without re-discovering the cluster.
func buildCodecsForGVKs(gvks gvkSupportedResources) (*serializer.CodecFactory, error) {
kubeScheme := runtime.NewScheme()
if err := registerDefaultKubeTypes(kubeScheme); err != nil {
return nil, trace.Wrap(err)
}
for _, gvk := range gvks {
if gvk == nil {
continue
}
// Skip well-known types already registered by registerDefaultKubeTypes.
if _, err := kubeScheme.New(*gvk); err == nil {
continue
}
kubeScheme.AddKnownTypeWithName(*gvk, &unstructured.Unstructured{})
kubeScheme.AddKnownTypeWithName(gvk.GroupVersion().WithKind(gvk.Kind+listSuffix), &unstructured.Unstructured{})
}
kubeCodecs := serializer.NewCodecFactory(kubeScheme)
return &kubeCodecs, nil
}
// getKubeAPIGroupAndVersion returns the API group and version from the given
// groupVersion string.
// The groupVersion string can be in the following formats:
// - "v1" -> group: "", version: "v1"
// - "<group>/<version>" -> group: "<group>", version: "<version>"
func getKubeAPIGroupAndVersion(groupVersion string) (group string, version string) {
splits := strings.Split(groupVersion, "/")
switch {
case len(splits) == 1:
return "", splits[0]
case len(splits) >= 2:
return splits[0], splits[1]
default:
return "", ""
}
}
/*
* 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 proxy
import (
"bytes"
"errors"
"fmt"
"io"
"net/http"
"path"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
authv1 "k8s.io/api/authorization/v1"
"k8s.io/apimachinery/pkg/runtime"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/utils"
)
// selfSubjectAccessReviews intercepts self subject access reviews requests and pre-validates
// them by applying the kubernetes resources RBAC rules to the request.
// If the self subject access review is allowed, the request is forwarded to the
// kubernetes API server for final validation.
func (f *Forwarder) selfSubjectAccessReviews(authCtx *authContext, w http.ResponseWriter, req *http.Request, p httprouter.Params) (resp any, err error) {
ctx, span := f.cfg.tracer.Start(
req.Context(),
"kube.Forwarder/selfSubjectAccessReviews",
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("selfSubjectAccessReviews"),
semconv.RPCSystemKey.String("kube"),
),
)
req = req.WithContext(ctx)
defer span.End()
sess, err := f.newClusterSession(req.Context(), *authCtx)
if err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(req.Context(), "Failed to create cluster session", "error", err)
return nil, trace.Wrap(err)
}
// sess.Close cancels the connection monitor context to release it sooner.
// When the server is under heavy load it can take a while to identify that
// the underlying connection is gone. This change prevents that and releases
// the resources as soon as we know the session is no longer active.
defer sess.close()
sess.upgradeToHTTP2 = true
sess.forwarder, err = f.makeSessionForwarder(sess)
if err != nil {
return nil, trace.Wrap(err)
}
// only allow self subject access reviews for the service that proxies the
// request to the kubernetes API server.
if sess.isLocalKubernetesCluster {
if err := f.validateSelfSubjectAccessReview(sess, w, req); trace.IsAccessDenied(err) {
return nil, nil
} else if err != nil {
return nil, trace.Wrap(err)
}
}
if err := f.setupForwardingHeaders(sess, req, true /* withImpersonationHeaders */); err != nil {
// This error goes to kubernetes client and is not visible in the logs
// of the teleport server if not logged here.
f.log.ErrorContext(req.Context(), "Failed to set up forwarding headers", "error", err)
return nil, trace.Wrap(err)
}
rw := httplib.NewResponseStatusRecorder(w)
sess.forwarder.ServeHTTP(rw, req)
f.emitAuditEvent(req, sess, rw.Status())
return nil, nil
}
// validateSelfSubjectAccessReview validates the self subject access review
// request by applying the kubernetes resources RBAC rules to the request.
func (f *Forwarder) validateSelfSubjectAccessReview(sess *clusterSession, w http.ResponseWriter, req *http.Request) error {
negotiator := newClientNegotiator(sess.codecFactory)
encoder, decoder, err := newEncoderAndDecoderForContentType(responsewriters.GetContentTypeHeader(req.Header), negotiator)
if err != nil {
return trace.Wrap(err)
}
accessReview, err := parseSelfSubjectAccessReviewRequest(decoder, req)
if err != nil {
return trace.Wrap(err)
}
// TODO(@creack): Remove this as part of the RBAC RFC. It grants excessive permissions.
if accessReview.Spec.ResourceAttributes == nil {
return nil
}
namespace := accessReview.Spec.ResourceAttributes.Namespace
group := accessReview.Spec.ResourceAttributes.Group
version := accessReview.Spec.ResourceAttributes.Version
resourceKind := accessReview.Spec.ResourceAttributes.Resource
resource, ok := sess.rbacSupportedResources.getResource(group, resourceKind)
if !ok {
// Mirror the data path so `kubectl auth can-i` matches a real request:
// discover the kind's group-version and, if it's still unknown, report denied.
details, derr := f.findKubeDetailsByClusterName(sess.kubeClusterName)
if derr != nil {
return nil
}
var found bool
resource, found = details.resolveResource(group, version, resourceKind)
if !found {
accessReview.Status = authv1.SubjectAccessReviewStatus{
Allowed: false,
Denied: true,
Reason: fmt.Sprintf(
"Kubernetes resource kind %q in API group %q is not known to this cluster",
resourceKind, group,
),
}
responsewriters.SetContentTypeHeader(w, req.Header)
if encodeErr := encoder.Encode(accessReview, w); encodeErr != nil {
return trace.Wrap(encodeErr)
}
// The denied response is already written; signal the caller to stop
// forwarding, like the other denial branches below.
return trace.AccessDenied("Kubernetes resource kind %q in API group %q is not known to this cluster", resourceKind, group)
}
}
name := accessReview.Spec.ResourceAttributes.Name
actx := sess.authContext
ident := actx.Identity.GetIdentity()
state, err := actx.CheckerContext.AccessStateFromTLSIdentity(req.Context(), &ident, f.cfg.CachingAuthClient)
if err != nil {
return trace.Wrap(err)
}
checker, err := actx.getAccessChecker()
if err != nil {
return trace.Wrap(err)
}
switch err := checker.Kube().CheckAccessToCluster(
actx.kubeCluster,
state,
services.RoleMatchers{
// Append a matcher that validates if the Kubernetes resource is allowed
// by the roles that satisfy the Kubernetes Cluster.
&kubernetesResourceMatcher{
resource: types.KubernetesResource{
Kind: resource.Name,
Name: name,
Namespace: namespace,
Verbs: []string{accessReview.Spec.ResourceAttributes.Verb},
APIGroup: accessReview.Spec.ResourceAttributes.Group,
},
isClusterWideResource: !resource.Namespaced,
},
}...); {
case errors.Is(err, services.ErrTrustedDeviceRequired):
return trace.Wrap(err)
case err != nil && resource.Namespaced:
namespaceNameToString := func(namespace, name string) string {
switch {
case namespace == "" && name == "":
return ""
case namespace != "" && name != "":
return path.Join(namespace, name)
case namespace != "":
return path.Join(namespace, "*")
default:
return path.Join("*", name)
}
}
accessReview.Status = authv1.SubjectAccessReviewStatus{
Allowed: false,
Denied: true,
Reason: fmt.Sprintf(
"access to %s %s denied by Teleport: please ensure that %q field in your Teleport role defines access to the desired resource.\n\n"+
"Valid example:\n"+
"kubernetes_resources:\n"+
"- kind: %s\n"+
" name: %s\n"+
" namespace: %s\n"+
" verbs: [%s]\n"+
" api_group: %s\n",
accessReview.Spec.ResourceAttributes.Resource,
namespaceNameToString(namespace, name),
kubernetesResourcesKey,
resource.Name,
emptyOrWildcard(name),
emptyOrWildcard(namespace),
emptyOrWildcard(""),
emptyOrWildcard(accessReview.Spec.ResourceAttributes.Group),
),
}
responsewriters.SetContentTypeHeader(w, req.Header)
if encodeErr := encoder.Encode(accessReview, w); encodeErr != nil {
return trace.Wrap(encodeErr)
}
return trace.Wrap(err)
case err != nil:
// If the request is for a cluster-wide resource, we need to grant access
// to it.
accessReview.Status = authv1.SubjectAccessReviewStatus{
Allowed: false,
Denied: true,
Reason: fmt.Sprintf(
"access to %s %s denied by Teleport: please ensure that %q field in your Teleport role defines access to the desired resource.\n\n"+
"Valid example:\n"+
"kubernetes_resources:\n"+
"- kind: %s\n"+
" name: %s\n"+
" verbs: [%s]\n"+
" api_group: %s",
accessReview.Spec.ResourceAttributes.Resource,
name,
kubernetesResourcesKey,
resource.Name,
emptyOrWildcard(name),
emptyOrWildcard(""),
emptyOrWildcard(accessReview.Spec.ResourceAttributes.Group),
),
}
responsewriters.SetContentTypeHeader(w, req.Header)
if encodeErr := encoder.Encode(accessReview, w); encodeErr != nil {
return trace.Wrap(encodeErr)
}
return trace.Wrap(err)
}
return nil
}
// emptyOrWildcard returns the string s if it is not empty, otherwise it returns
// '*'.
func emptyOrWildcard(s string) string {
if s == "" {
return fmt.Sprintf("'%s'", types.Wildcard)
}
return s
}
// parseSelfSubjectAccessReviewRequest parses the request body into a SelfSubjectAccessReview object
// and replaces the body so it can be read again.
func parseSelfSubjectAccessReviewRequest(decoder runtime.Decoder, req *http.Request) (*authv1.SelfSubjectAccessReview, error) {
payload, err := io.ReadAll(req.Body)
if err != nil {
return nil, trace.Wrap(err)
}
req.Body.Close()
req.Body = io.NopCloser(bytes.NewReader(payload))
gvk := authv1.SchemeGroupVersion.WithKind("SelfSubjectAccessReview")
obj, err := decodeAndSetGVK(decoder, payload, &gvk)
if err != nil {
return nil, trace.Wrap(err)
}
switch o := obj.(type) {
case *authv1.SelfSubjectAccessReview:
return o, nil
default:
return nil, trace.BadParameter("unexpected object type: %T", obj)
}
}
// kubernetesResourceMatcher matches a role against a Kubernetes Resource.
// This matcher is different form services.KubernetesResourceMatcher because
// if skips some validations if the user doesn't ask for a specific resource.
// If name and namespace are empty, it means that the user wants to match all
// resources of the specified kind for which the user might have access to.
// If the user asks for name="", namespace="" and the role has a matcher
// with name="foo", namespace="bar", the matcher will match but the user
// might not be able to see any resource if the resource does not exist
// in the cluster.
// This matcher assumes the role's kubernetes_resources configured eventually
// match with resources that exist in the cluster.
type kubernetesResourceMatcher struct {
resource types.KubernetesResource
isClusterWideResource bool
}
// Match matches a Kubernetes Resource against provided role and condition.
func (m *kubernetesResourceMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
resources := role.GetKubeResources(condition)
if len(resources) == 0 {
return false, nil
}
kind := m.resource.Kind
name := m.resource.Name
namespace := m.resource.Namespace
// If the resource is global, clear the namespace.
// NOTE: kubectl will yield a warning for this case, but we still need to process the request.
if m.isClusterWideResource {
namespace = ""
}
// If we are dealing with a namespace resource, consider it cluster wide even though it is not.
if m.resource.Kind == "namespaces" {
m.isClusterWideResource = true
}
for _, resource := range resources {
isResourceTheSameKind := kind == resource.Kind || resource.Kind == types.Wildcard
namespaceScopeMatch := resource.Kind == "namespaces" && !m.isClusterWideResource
if !isResourceTheSameKind && !namespaceScopeMatch {
continue
}
if len(m.resource.Verbs) == 1 && !utils.IsVerbAllowed(resource.Verbs, m.resource.Verbs[0]) {
continue
}
switch ok, err := utils.SliceMatchesRegex(m.resource.APIGroup, []string{resource.APIGroup}); {
case err != nil:
return false, trace.Wrap(err)
case !ok:
continue
}
// If the resource name and namespace are empty, it means that the
// user wants to match all resources of the specified kind.
// We can return true immediately because the user is allowed to get resources
// of the specified kind but might not be able to see any if the matchers do not
// match with any resource.
if (resource.Namespace == "" || resource.Namespace == types.Wildcard) && name == "" && namespace == "" {
return true, nil
}
// If the resource name isn't empty but the resource kind is a namespace scope
// match - i.e. the resource.Kind==types.KindKubeNamespace and the desired
// resource kind is not a cluster-wide resource - we should skip the resource
// name validation.
if name != "" && !namespaceScopeMatch {
switch ok, err := utils.SliceMatchesRegex(name, []string{resource.Name}); {
case err != nil:
return false, trace.Wrap(err)
case !ok:
continue
}
}
if resource.Kind == "namespaces" && namespace != "" {
if ok, err := utils.SliceMatchesRegex(namespace, []string{resource.Name}); err != nil || ok {
return ok, trace.Wrap(err)
}
} else {
if ok, err := utils.SliceMatchesRegex(namespace, []string{resource.Namespace}); err != nil || ok {
return ok, trace.Wrap(err)
}
}
}
return false, nil
}
/*
* 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 proxy
import (
"context"
"crypto/tls"
"log/slog"
"maps"
"net"
"net/http"
"sync"
"time"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/service/eks"
"github.com/aws/aws-sdk-go-v2/service/sts"
"github.com/gravitational/trace"
"golang.org/x/net/http2"
"github.com/gravitational/teleport"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/cloud/azure"
"github.com/gravitational/teleport/lib/cloud/gcp"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/healthcheck"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/inventory"
kubewatcher "github.com/gravitational/teleport/lib/kube/proxy/watcher"
"github.com/gravitational/teleport/lib/labels"
"github.com/gravitational/teleport/lib/limiter"
"github.com/gravitational/teleport/lib/multiplexer"
"github.com/gravitational/teleport/lib/relaytunnel"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/srv/ingress"
"github.com/gravitational/teleport/lib/utils/aws/stsutils"
"github.com/gravitational/teleport/lib/utils/log"
)
// TLSServerConfig is a configuration for TLS server
type TLSServerConfig struct {
// ForwarderConfig is a config of a forwarder
ForwarderConfig
// TLS is a base TLS configuration
TLS *tls.Config
// LimiterConfig is limiter config
LimiterConfig limiter.Config
// AccessPoint is caching access point
AccessPoint authclient.ReadKubernetesAccessPoint
// OnHeartbeat is a callback for kubernetes_service heartbeats.
OnHeartbeat func(error)
// GetRotation returns the certificate rotation state.
GetRotation services.RotationGetter
// ConnectedProxyGetter gets the proxies teleport is connected to.
ConnectedProxyGetter reversetunnelclient.ConnectedProxyGetter
// RelayInfoGetter is the function used to get the relay tunnel client info
// to fill in the server heartbeats.
RelayInfoGetter relaytunnel.GetRelayInfoFunc
// Log is the logger.
Log *slog.Logger
// Selectors is a list of resource monitor selectors.
ResourceMatchers []services.ResourceMatcher
// OnReconcile is called after each kube_cluster resource reconciliation.
OnReconcile func(types.KubeClusters)
// azureClients provides Azure SDK clients
azureClients azure.Clients
// gcpClients provides GCP SDK clients
gcpClients gcp.Clients
// awsCloudClients provides AWS SDK clients.
awsClients *awsClientsGetter
// StaticLabels is a map of static labels associated with this service.
// Each cluster advertised by this kubernetes_service will include these static labels.
// If the service and a cluster define labels with the same key,
// service labels take precedence over cluster labels.
// Used for RBAC.
StaticLabels map[string]string
// DynamicLabels define the dynamic labels associated with this service.
// Each cluster advertised by this kubernetes_service will include these dynamic labels.
// If the service and a cluster define labels with the same key,
// service labels take precedence over cluster labels.
// Used for RBAC.
DynamicLabels *labels.Dynamic
// CloudLabels is a map of static labels imported from a cloud provider associated with this
// service. Used for RBAC.
// If StaticLabels and CloudLabels define labels with the same key,
// StaticLabels take precedence over CloudLabels.
CloudLabels labels.Importer
// IngressReporter reports new and active connections.
IngressReporter *ingress.Reporter
// KubernetesServersWatcher is used by the kube proxy to watch for changes in the
// kubernetes servers of a cluster. Proxy requires it to update the kubeServersMap
// which holds the list of kubernetes_services connected to the proxy for a given
// kubernetes cluster name. Proxy uses this map to route requests to the correct
// kubernetes_service. The servers are kept in memory to avoid making unnecessary
// unmarshal calls followed by filtering and to improve memory usage.
KubernetesServersWatcher *kubewatcher.ProxyKubeServerWatcher
// PROXYProtocolMode controls behavior related to unsigned PROXY protocol headers.
PROXYProtocolMode multiplexer.PROXYProtocolMode
// InventoryHandle is used to send kube server heartbeats via the inventory control stream.
InventoryHandle inventory.DownstreamHandle
// HealthCheckManager manages checking the health of Kubernetes clusters.
HealthCheckManager healthcheck.Manager
}
type awsClientsGetter struct{}
func (f *awsClientsGetter) GetConfig(ctx context.Context, region string, optFns ...awsconfig.OptionsFn) (aws.Config, error) {
return awsconfig.GetConfig(ctx, region, optFns...)
}
func (f *awsClientsGetter) GetAWSEKSClient(cfg aws.Config) EKSClient {
return eks.NewFromConfig(cfg)
}
func (f *awsClientsGetter) GetAWSSTSPresignClient(cfg aws.Config) STSPresignClient {
stsClient := stsutils.NewFromConfig(cfg)
return sts.NewPresignClient(stsClient)
}
// CheckAndSetDefaults checks and sets default values
func (c *TLSServerConfig) CheckAndSetDefaults() error {
if err := c.ForwarderConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if c.TLS == nil {
return trace.BadParameter("missing parameter TLS")
}
if c.AccessPoint == nil {
return trace.BadParameter("missing parameter AccessPoint")
}
if c.InventoryHandle == nil {
return trace.BadParameter("missing parameter InventoryHandle")
}
if c.ConnectedProxyGetter == nil {
return trace.BadParameter("missing parameter ConnectedProxyGetter")
}
if c.HealthCheckManager == nil {
return trace.BadParameter("missing parameter HealthCheckManager")
}
if err := c.validateLabelKeys(); err != nil {
return trace.Wrap(err)
}
switch c.KubeServiceType {
case ProxyService, LegacyProxyService:
if c.KubernetesServersWatcher == nil {
return trace.BadParameter("missing parameter KubernetesServersWatcher")
}
case KubeService:
if c.GetScope() != "" && c.KubernetesServersWatcher != nil {
return trace.BadParameter("KubernetesServersWatcher is not supported for scoped KubeService")
}
}
if c.Log == nil {
c.Log = slog.Default()
}
if c.azureClients == nil {
azureClients, err := azure.NewClients()
if err != nil {
return trace.Wrap(err)
}
c.azureClients = azureClients
}
if c.gcpClients == nil {
c.gcpClients = gcp.NewClients()
}
if c.awsClients == nil {
c.awsClients = &awsClientsGetter{}
}
return nil
}
// validateLabelKeys checks that all labels keys are valid.
// Dynamic labels are validated in labels.NewDynamicLabels.
func (c *TLSServerConfig) validateLabelKeys() error {
for name := range c.StaticLabels {
if !types.IsValidLabelKey(name) {
return trace.BadParameter("invalid label key: %q", name)
}
}
return nil
}
// TLSServer is TLS auth server
type TLSServer struct {
*http.Server
// TLSServerConfig is TLS server configuration used for auth server
TLSServerConfig
fwd *Forwarder
mu sync.Mutex
listener net.Listener
heartbeats map[string]*srv.HeartbeatV2
closeContext context.Context
closeFunc context.CancelFunc
// kubeClusterWatcher monitors changes to kube cluster resources.
kubeClusterWatcher *services.GenericWatcher[types.KubeCluster, readonly.KubeCluster]
// reconciler reconciles proxied kube clusters with kube_clusters resources.
reconciler *services.Reconciler[types.KubeCluster]
// monitoredKubeClusters contains all kube clusters the proxied kube_clusters are
// reconciled against.
monitoredKubeClusters monitoredKubeClusters
// reconcileCh triggers reconciliation of proxied kube_clusters.
reconcileCh chan struct{}
log *slog.Logger
}
// NewTLSServer returns new unstarted TLS server
func NewTLSServer(cfg TLSServerConfig) (*TLSServer, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
log := cfg.Log.With(teleport.ComponentKey, cfg.Component)
// limiter limits requests by frequency and amount of simultaneous
// connections per client
limiter, err := limiter.NewLimiter(cfg.LimiterConfig)
if err != nil {
return nil, trace.Wrap(err)
}
cfg.ForwarderConfig.log = log
fwd, err := NewForwarder(cfg.ForwarderConfig)
if err != nil {
return nil, trace.Wrap(err)
}
if len(fwd.kubeClusters()) == 0 && cfg.KubeServiceType == KubeService &&
len(cfg.ResourceMatchers) == 0 {
// if fwd has no clusters and the service type is KubeService but no resource watcher is configured
// then the kube_service does not need to start since it will not serve any static or dynamic cluster.
return nil, trace.BadParameter("kube_service won't start because it has neither static clusters nor a resource watcher configured.")
}
clustername, err := cfg.AccessPoint.GetClusterName(cfg.Context)
if err != nil {
return nil, trace.Wrap(err)
}
// authMiddleware authenticates request assuming TLS client authentication
// adds authentication information to the context
// and passes it to the API server
authMiddleware := &authz.Middleware{
ClusterName: clustername.GetClusterName(),
AcceptedUsage: []string{teleport.UsageKubeOnly},
// EnableCredentialsForwarding is set to true to allow the proxy to forward
// the client identity to the target service using headers instead of TLS
// certificates. This is required for the kube service and leaf cluster proxy
// to be able to replace the client identity with the header payload when
// the request is forwarded from a Teleport Proxy.
EnableCredentialsForwarding: true,
Handler: fwd,
}
// Wrap sets the next middleware in chain to the authMiddleware
limiter.WrapHandle(authMiddleware)
// force client auth if given
cfg.TLS.ClientAuth = tls.VerifyClientCertIfGiven
tracingHandler := httplib.MakeTracingHandler(limiter, teleport.ComponentKube)
kubeHTTPserver := newKubeHTTPServer(tracingHandler, log, cfg.TLS, cfg.IngressReporter)
server := &TLSServer{
fwd: fwd,
TLSServerConfig: cfg,
Server: kubeHTTPserver,
heartbeats: make(map[string]*srv.HeartbeatV2),
monitoredKubeClusters: monitoredKubeClusters{
static: fwd.kubeClusters(),
},
reconcileCh: make(chan struct{}),
log: log,
}
server.TLS.GetConfigForClient = server.GetConfigForClient
server.closeContext, server.closeFunc = context.WithCancel(cfg.Context)
// register into the forwarder the method to get kubernetes servers for a kube cluster.
server.fwd.getKubernetesServersForKubeCluster, err = server.getKubernetesServersForKubeClusterFunc()
if err != nil {
return nil, trace.Wrap(err)
}
return server, nil
}
// newKubeHTTPServer builds the kube proxy's *http.Server.
func newKubeHTTPServer(inner http.Handler, log *slog.Logger, tlsCfg *tls.Config, reporter *ingress.Reporter) *http.Server {
return &http.Server{
Handler: newTimeoutResetHandler(inner, log),
ReadHeaderTimeout: apidefaults.DefaultIOTimeout * 2,
// WriteTimeout drives the TLS handshake deadline via net/http's
// tlsHandshakeTimeout = min(ReadHeaderTimeout, ReadTimeout, WriteTimeout) formula,
// bounding the pre-authentication goroutine hold time.
WriteTimeout: defaults.HandshakeReadDeadline,
IdleTimeout: apidefaults.DefaultIdleTimeout,
TLSConfig: tlsCfg,
ConnState: ingress.HTTPConnStateReporter(ingress.Kube, reporter),
ConnContext: func(ctx context.Context, c net.Conn) context.Context {
return authz.ContextWithClientAddrs(ctx, c.RemoteAddr(), c.LocalAddr())
},
}
}
// newTimeoutResetHandler wraps next so that the per-request write deadline
// (set by net/http from Server.WriteTimeout) is cleared before next runs.
func newTimeoutResetHandler(next http.Handler, log *slog.Logger) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if err := http.NewResponseController(w).SetWriteDeadline(time.Time{}); err != nil {
log.ErrorContext(r.Context(), "failed to reset response write deadline", "error", err)
}
next.ServeHTTP(w, r)
})
}
// ServeOption is a functional option for the multiplexer.
type ServeOption func(*multiplexer.Config)
// WithMultiplexerIgnoreSelfConnections is used for tests, it makes multiplexer ignore the fact that it's self
// connection (coming from same IP as the listening address) when deciding if it should drop connection with
// missing required PROXY header. This is needed since all connections in tests are self connections.
func WithMultiplexerIgnoreSelfConnections() ServeOption {
return func(cfg *multiplexer.Config) {
cfg.IgnoreSelfConnections = true
}
}
// Serve takes TCP listener, upgrades to TLS using config and starts serving
func (t *TLSServer) Serve(listener net.Listener, options ...ServeOption) error {
caGetter := func(ctx context.Context, id types.CertAuthID, loadKeys bool) (types.CertAuthority, error) {
return t.CachingAuthClient.GetCertAuthority(ctx, id, loadKeys)
}
muxConfig := multiplexer.Config{
Context: t.Context,
Listener: listener,
Clock: t.Clock,
PROXYProtocolMode: t.PROXYProtocolMode,
ID: t.Component,
CertAuthorityGetter: caGetter,
LocalClusterName: t.ClusterName,
// Increases deadline until the agent receives the first byte to 10s.
// It's required to accommodate setups with high latency and where the time
// between the TCP being accepted and the time for the first byte is longer
// than the default value - 1s.
DetectTimeout: 10 * time.Second,
}
for _, opt := range options {
opt(&muxConfig)
}
// Wrap listener with a multiplexer to get PROXY Protocol support.
mux, err := multiplexer.New(muxConfig)
if err != nil {
return trace.Wrap(err)
}
go mux.Serve()
defer mux.Close()
t.mu.Lock()
select {
// If the server is closed before the listener is started, return early
// to avoid deadlock.
case <-t.closeContext.Done():
t.mu.Unlock()
return nil
default:
}
t.listener = mux.TLS()
err = http2.ConfigureServer(t.Server, &http2.Server{})
t.mu.Unlock()
if err != nil {
return trace.Wrap(err)
}
// startStaticClusterHeartbeats starts the heartbeat process for static clusters.
// static clusters can be specified via kubeconfig or clusterName for Teleport agent
// running in Kubernetes.
if err := t.startStaticClustersHeartbeat(); err != nil {
return trace.Wrap(err)
}
// Start reconciler that will be reconciling proxied clusters with
// kube_cluster resources.
if err := t.startReconciler(t.closeContext); err != nil {
return trace.Wrap(err)
}
// Initialize watcher that will be dynamically (un-)registering
// proxied clusters based on the kube_cluster resources.
// This watcher is only started for the kube_service if a resource watcher
// is configured.
kubeClusterWatcher, err := t.startKubeClusterResourceWatcher(t.closeContext)
if err != nil {
return trace.Wrap(err)
}
t.mu.Lock()
t.kubeClusterWatcher = kubeClusterWatcher
t.mu.Unlock()
if t.OnHeartbeat != nil {
// Kube uses heartbeat v2, which heartbeats resources but not the server itself
// If there are no resources, we will never report ready.
// We work around by reporting ready after the first successful watcher init.
var watcherWaiters []func() error
if kubeClusterWatcher != nil {
watcherWaiters = append(watcherWaiters, kubeClusterWatcher.WaitInitialization)
}
if t.KubernetesServersWatcher != nil {
watcherWaiters = append(watcherWaiters, t.KubernetesServersWatcher.WaitInitialization)
}
if len(watcherWaiters) > 0 {
go func() {
for _, w := range watcherWaiters {
err := w()
if err != nil {
t.OnHeartbeat(err)
return
}
}
t.OnHeartbeat(nil)
}()
}
}
// kubeServerWatcher is used by the kube proxy to watch for changes in the
// kubernetes servers of a cluster. Proxy requires it to update the kubeServersMap
// which holds the list of kubernetes_services connected to the proxy for a given
// kubernetes cluster name. Proxy uses this map to route requests to the correct
// kubernetes_service. The servers are kept in memory to avoid making unnecessary
// unmarshal calls followed by filtering to improve memory usage.
if t.KubernetesServersWatcher != nil {
// Wait for the watcher to initialize before starting the server so that the
// proxy can start routing requests to the kubernetes_service instead of
// returning an error because the cache is not initialized.
if err := t.KubernetesServersWatcher.WaitInitialization(); err != nil {
return trace.Wrap(err)
}
}
return t.Server.Serve(tls.NewListener(mux.TLS(), t.TLS))
}
// Close closes the server and cleans up all resources.
func (t *TLSServer) Close() error {
return trace.Wrap(t.close(t.closeContext))
}
// Shutdown closes the server and cleans up all resources.
func (t *TLSServer) Shutdown(ctx context.Context) error {
// TODO(tigrato): handle connections gracefully and wait for them to finish.
// This might be problematic because exec and port forwarding connections
// are long lived connections and if we wait for them to finish, we might
// end up waiting forever.
return trace.Wrap(t.close(ctx))
}
// close closes the server and cleans up all resources.
func (t *TLSServer) close(ctx context.Context) error {
var errs []error
for _, kubeCluster := range t.fwd.kubeClusters() {
errs = append(errs, t.unregisterKubeCluster(ctx, kubeCluster, true))
}
errs = append(errs, t.fwd.Close(), t.Server.Close())
t.closeFunc()
t.mu.Lock()
kubeClusterWatcher := t.kubeClusterWatcher
t.mu.Unlock()
// Stop the kube_cluster resource watcher.
if kubeClusterWatcher != nil {
kubeClusterWatcher.Close()
}
// Stop the kube_server resource watcher.
if t.KubernetesServersWatcher != nil {
t.KubernetesServersWatcher.Close()
}
var listClose error
t.mu.Lock()
if t.listener != nil {
listClose = t.listener.Close()
}
t.mu.Unlock()
errs = append(errs, t.gcpClients.Close())
return trace.NewAggregate(append(errs, listClose)...)
}
// GetConfigForClient is getting called on every connection
// and server's GetConfigForClient reloads the list of trusted
// local and remote certificate authorities
func (t *TLSServer) GetConfigForClient(info *tls.ClientHelloInfo) (*tls.Config, error) {
return authclient.WithClusterCAs(t.TLS, t.AccessPoint, t.ClusterName, t.log)(info)
}
// GetServerInfo returns a services.Server object for heartbeats (aka
// presence).
func (t *TLSServer) GetServerInfo(name string) (*types.KubernetesServerV3, error) {
t.mu.Lock()
defer t.mu.Unlock()
var addr string
if t.TLSServerConfig.ForwarderConfig.PublicAddr != "" {
addr = t.TLSServerConfig.ForwarderConfig.PublicAddr
} else if t.listener != nil {
addr = t.listener.Addr().String()
}
cluster, err := t.getKubeClusterWithServiceLabels(name)
if err != nil {
return nil, trace.Wrap(err)
}
// Both proxy and kubernetes services can run in the same instance (same
// cluster names). Add a name suffix to make them distinct.
//
// Note: we *don't* want to add suffix for kubernetes_service!
// This breaks reverse tunnel routing, which uses server.Name.
if t.KubeServiceType != KubeService {
name += teleport.KubeLegacyProxySuffix
}
var relayGroup string
var relayIDs []string
if t.RelayInfoGetter != nil {
// relayInfoGetter returns a copy of the slice, so we can move it in the
// protobuf message
relayGroup, relayIDs = t.RelayInfoGetter()
}
srv, err := types.NewKubernetesServerV3(
types.Metadata{
Name: name,
Namespace: t.Namespace,
},
types.KubernetesServerSpecV3{
Version: teleport.Version,
Hostname: addr,
HostID: t.TLSServerConfig.HostID,
Rotation: t.getRotationState(),
Cluster: cluster,
ProxyIDs: t.ConnectedProxyGetter.GetProxyIDs(),
RelayGroup: relayGroup,
RelayIds: relayIDs,
},
// getKubeClusterWithServiceLabels already ensures that the cluster has the correct scope and that scope
// is usable by this forwarder. We only need to make sure that the kube server we build shares the same scope
types.KubeServerWithScope(cluster.GetScope()),
)
if err != nil {
return nil, trace.Wrap(err)
}
srv.SetExpiry(t.Clock.Now().UTC().Add(apidefaults.ServerAnnounceTTL))
// Get the kube cluster health and send it to the auth server.
srv.SetTargetHealth(t.getTargetHealth(t.closeContext, cluster))
return srv, nil
}
// startHealthCheck starts checking the health of a Kubernetes cluster.
func (t *TLSServer) startHealthCheck(cluster types.KubeCluster) error {
kubeDetails, err := t.fwd.findKubeDetailsByClusterName(cluster.GetName())
if err != nil {
return trace.Wrap(err)
}
err = t.HealthCheckManager.AddTarget(healthcheck.Target{
HealthChecker: kubeDetails,
GetResource: func() types.ResourceWithLabels { return cluster },
})
return trace.Wrap(err)
}
// stopHealthCheck stops checking the health of a Kubernetes cluster.
func (t *TLSServer) stopHealthCheck(cluster types.KubeCluster) error {
if err := t.HealthCheckManager.RemoveTarget(cluster); err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err)
}
return nil
}
// startHeartbeatAndHealthCheck starts heart beats and health checks.
func (t *TLSServer) startHeartbeatAndHealthCheck(cluster types.KubeCluster) error {
if err := t.startHealthCheck(cluster); err != nil {
return trace.Wrap(err)
}
if err := t.startHeartbeat(cluster.GetName()); err != nil {
return trace.Wrap(err)
}
return nil
}
// stopHeartbeatAndHealthCheck stops heart beats and health checks.
func (t *TLSServer) stopHeartbeatAndHealthCheck(cluster types.KubeCluster) error {
var errs []error
if err := t.stopHealthCheck(cluster); err != nil {
errs = append(errs, err)
}
if err := t.stopHeartbeat(cluster.GetName()); err != nil {
errs = append(errs, err)
}
return trace.NewAggregate(errs...)
}
// getTargetHealth returns the health of a Kubernetes cluster.
func (t *TLSServer) getTargetHealth(ctx context.Context, cluster types.KubeCluster) *types.TargetHealth {
health, err := t.HealthCheckManager.GetTargetHealth(cluster)
if err == nil {
return health
}
if trace.IsNotFound(err) {
return &types.TargetHealth{
Status: string(types.TargetHealthStatusUnknown),
TransitionReason: string(types.TargetHealthTransitionReasonDisabled),
Message: "Unable to find the Kubernetes cluster",
}
}
t.log.WarnContext(ctx, "Failed to get kube cluster health",
"kube_cluster", log.StringerAttr(cluster),
"error", err,
)
return &types.TargetHealth{
Status: string(types.TargetHealthStatusUnknown),
TransitionReason: string(types.TargetHealthTransitionReasonInternalError),
TransitionError: err.Error(),
Message: "Teleport failed to get the Kubernetes cluster health status (this is a bug)",
}
}
// getKubeClusterWithServiceLabels finds the kube cluster by name, strips the credentials,
// replaces the cluster dynamic labels with their latest value available and updates
// the cluster with the service dynamic and static labels.
// We strip the Azure, AWS and Kubeconfig credentials so they are not leaked when
// heartbeating the cluster.
func (t *TLSServer) getKubeClusterWithServiceLabels(name string) (*types.KubernetesClusterV3, error) {
// it is safe do read from details since the structure is never updated.
// we replace the whole structure each time an update happens to a dynamic cluster.
details, err := t.fwd.findKubeDetailsByClusterName(name)
if err != nil {
return nil, trace.Wrap(err)
}
// NewKubernetesClusterV3WithoutSecrets creates a copy of details.kubeCluster without
// any credentials or cloud access details.
clusterWithoutCreds, err := types.NewKubernetesClusterV3WithoutSecrets(details.kubeCluster)
if err != nil {
return nil, trace.Wrap(err)
}
// The Proxy Service forwarder will always be unscoped and needs to be able to forward
// to scoped clusters as well.
if t.GetScope() != "" {
if scopes.Compare(t.GetScope(), details.kubeCluster.GetScope()) != scopes.Equivalent {
// This should only happen if there's a bug in scoped access checking for KubernetesCluster resources.
// The kube proxy should never have access to clusters from orthogonal scopes. We also block access
// to clusters in child scopes but this may be relaxed in the future.
return nil, trace.AccessDenied("kube forwarder found kube cluster from different scope")
}
}
if details.dynamicLabels != nil {
clusterWithoutCreds.SetDynamicLabels(details.dynamicLabels.Get())
}
t.setServiceLabels(clusterWithoutCreds)
return clusterWithoutCreds, nil
}
// startHeartbeat starts the registration heartbeat to the auth server.
func (t *TLSServer) startHeartbeat(name string) error {
heartbeat, err := srv.NewKubernetesServerHeartbeat(srv.HeartbeatV2Config[*types.KubernetesServerV3]{
InventoryHandle: t.InventoryHandle,
GetResource: func(context.Context) (*types.KubernetesServerV3, error) { return t.GetServerInfo(name) },
OnHeartbeat: t.TLSServerConfig.OnHeartbeat,
})
if err != nil {
return trace.Wrap(err)
}
go heartbeat.Run()
t.mu.Lock()
defer t.mu.Unlock()
t.heartbeats[name] = heartbeat
return nil
}
// getRotationState is a helper to return this server's CA rotation state.
func (t *TLSServer) getRotationState() types.Rotation {
rotation, err := t.TLSServerConfig.GetRotation(types.RoleKube)
if err != nil && !trace.IsNotFound(err) {
t.log.WarnContext(t.closeContext, "Failed to get rotation state", "error", err)
}
if rotation != nil {
return *rotation
}
return types.Rotation{}
}
func (t *TLSServer) startStaticClustersHeartbeat() error {
// Start the heartbeat to announce kubernetes_service presence.
//
// Only announce when running in an actual kube_server, or when
// running in proxy_service with local kube credentials. This means that
// proxy_service will pretend to also be kube_server.
if t.KubeServiceType == KubeService ||
t.KubeServiceType == LegacyProxyService {
t.log.DebugContext(t.closeContext, "Starting kubernetes_service heartbeats and health checks")
for _, kc := range t.fwd.kubeClusters() {
if err := t.startHeartbeatAndHealthCheck(kc); err != nil {
return trace.Wrap(err)
}
}
} else {
t.log.DebugContext(t.closeContext, "No local kube credentials on proxy, will not start kubernetes_service heartbeats and health checks")
}
return nil
}
// stopHeartbeat stops the registration heartbeat to the auth server.
func (t *TLSServer) stopHeartbeat(name string) error {
t.mu.Lock()
defer t.mu.Unlock()
heartbeat, ok := t.heartbeats[name]
if !ok {
return nil
}
delete(t.heartbeats, name)
return trace.Wrap(heartbeat.Close())
}
// getServiceStaticLabels gets the labels that the server should present as static,
// which includes Cloud labels if available.
func (t *TLSServer) getServiceStaticLabels() map[string]string {
if t.CloudLabels == nil {
return t.StaticLabels
}
labels := maps.Clone(t.CloudLabels.Get())
// Let static labels override ec2 labels.
maps.Copy(labels, t.StaticLabels)
return labels
}
// setServiceLabels updates the cluster labels with the kubernetes_service labels.
// If the cluster and the service define overlapping labels the service labels take precedence.
// This function manipulates the original cluster.
func (t *TLSServer) setServiceLabels(cluster types.KubeCluster) {
serviceStaticLabels := t.getServiceStaticLabels()
if len(serviceStaticLabels) > 0 {
staticLabels := cluster.GetStaticLabels()
if staticLabels == nil {
staticLabels = make(map[string]string)
}
// if cluster and service define the same static label key, service labels have precedence.
maps.Copy(staticLabels, serviceStaticLabels)
cluster.SetStaticLabels(staticLabels)
}
if t.DynamicLabels != nil {
dstDynLabels := cluster.GetDynamicLabels()
if dstDynLabels == nil {
dstDynLabels = map[string]types.CommandLabel{}
}
// get service level dynamic labels.
serviceDynLabels := t.DynamicLabels.Get()
// if cluster and service define the same dynamic label key, service labels have precedence.
maps.Copy(dstDynLabels, serviceDynLabels)
cluster.SetDynamicLabels(dstDynLabels)
}
}
// getKubernetesServersForKubeClusterFunc returns a function that returns the kubernetes servers
// for a given kube cluster depending on the type of service.
func (t *TLSServer) getKubernetesServersForKubeClusterFunc() (getKubeServersByNameFunc, error) {
switch t.KubeServiceType {
case KubeService:
return func(_ context.Context, name string) ([]types.KubeServer, error) {
// If this is a kube_service, we can just return the local kube servers.
kube, err := t.getKubeClusterWithServiceLabels(name)
if err != nil {
return nil, trace.Wrap(err)
}
srv, err := types.NewKubernetesServerV3FromCluster(kube, "", t.HostID)
if err != nil {
return nil, trace.Wrap(err)
}
return []types.KubeServer{srv}, nil
}, nil
case ProxyService:
return t.KubernetesServersWatcher.GetKubeServersForClusterName, nil
case LegacyProxyService:
return func(ctx context.Context, name string) ([]types.KubeServer, error) {
// If this is a legacy kube proxy, then we need to return the local kube servers if
// the local server is proxying the target cluster, otherwise act like a proxy_service.
// and forward the request to the next proxy.
kube, err := t.getKubeClusterWithServiceLabels(name)
if err != nil {
servers, err := t.KubernetesServersWatcher.GetKubeServersForClusterName(ctx, name)
return servers, trace.Wrap(err)
}
srv, err := types.NewKubernetesServerV3FromCluster(kube, "", t.HostID)
if err != nil {
return nil, trace.Wrap(err)
}
return []types.KubeServer{srv}, nil
}, nil
default:
return nil, trace.BadParameter("unknown kubernetes service type %q", t.KubeServiceType)
}
}
/*
* 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 proxy
import (
"context"
"fmt"
"io"
"log/slog"
"net/http"
"path"
"slices"
"strings"
"sync"
"time"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
corev1 "k8s.io/api/core/v1"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/fields"
"k8s.io/apimachinery/pkg/runtime"
apimachinerytypes "k8s.io/apimachinery/pkg/types"
"k8s.io/apimachinery/pkg/watch"
"k8s.io/client-go/tools/cache"
"k8s.io/client-go/tools/remotecommand"
watchtools "k8s.io/client-go/tools/watch"
"github.com/gravitational/teleport"
kubewaitingcontainerpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/kubewaitingcontainer/v1"
"github.com/gravitational/teleport/api/types"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/auth/moderation"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/events/recorder"
"github.com/gravitational/teleport/lib/kube/proxy/streamproto"
"github.com/gravitational/teleport/lib/services"
tsession "github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
const sessionRecorderID = "session-recorder"
const (
PresenceVerifyInterval = time.Second * 15
PresenceMaxDifference = time.Minute
sessionMaxLifetime = time.Hour * 24
)
// remoteClient is either a kubectl or websocket client.
type remoteClient interface {
queueID() uuid.UUID
stdinStream() io.Reader
stdoutStream() io.Writer
stderrStream() io.Writer
resizeQueue() <-chan terminalResizeMessage
resize(size *remotecommand.TerminalSize) error
forceTerminate() <-chan struct{}
sendStatus(error) error
io.Closer
}
type websocketClientStreams struct {
id uuid.UUID
stream *streamproto.SessionStream
}
func (p *websocketClientStreams) queueID() uuid.UUID {
return p.id
}
func (p *websocketClientStreams) stdinStream() io.Reader {
return p.stream
}
func (p *websocketClientStreams) stdoutStream() io.Writer {
return p.stream
}
func (p *websocketClientStreams) stderrStream() io.Writer {
return p.stream
}
func (p *websocketClientStreams) resizeQueue() <-chan terminalResizeMessage {
ch := make(chan terminalResizeMessage)
go func() {
defer close(ch)
for {
select {
case <-p.stream.Done():
return
case size := <-p.stream.ResizeQueue():
if size == nil {
return
}
ch <- terminalResizeMessage{
size: size,
source: p.id,
}
}
}
}()
return ch
}
func (p *websocketClientStreams) resize(size *remotecommand.TerminalSize) error {
return p.stream.Resize(size)
}
func (p *websocketClientStreams) forceTerminate() <-chan struct{} {
return p.stream.ForceTerminateQueue()
}
func (p *websocketClientStreams) sendStatus(err error) error {
return nil
}
func (p *websocketClientStreams) Close() error {
return trace.Wrap(p.stream.Close())
}
type kubeProxyClientStreams struct {
id uuid.UUID
proxy *remoteCommandProxy
sizeQueue *termQueue
stdin io.Reader
stdout io.Writer
stderr io.Writer
close chan struct{}
wg sync.WaitGroup
}
func newKubeProxyClientStreams(proxy *remoteCommandProxy) *kubeProxyClientStreams {
options := proxy.options()
return &kubeProxyClientStreams{
id: uuid.New(),
proxy: proxy,
stdin: options.Stdin,
stdout: options.Stdout,
stderr: options.Stderr,
close: make(chan struct{}),
sizeQueue: proxy.resizeQueue,
}
}
func (p *kubeProxyClientStreams) queueID() uuid.UUID {
return p.id
}
func (p *kubeProxyClientStreams) stdinStream() io.Reader {
return p.stdin
}
func (p *kubeProxyClientStreams) stdoutStream() io.Writer {
return p.stdout
}
func (p *kubeProxyClientStreams) stderrStream() io.Writer {
return p.stderr
}
func (p *kubeProxyClientStreams) resizeQueue() <-chan terminalResizeMessage {
ch := make(chan terminalResizeMessage)
if p.sizeQueue == nil {
return ch
}
p.wg.Add(1)
go func() {
defer p.wg.Done()
for {
size := p.sizeQueue.Next()
if size == nil {
return
}
select {
case ch <- terminalResizeMessage{size, p.id}:
// Check if the sizeQueue was already terminated.
case <-p.sizeQueue.done.Done():
return
}
}
}()
return ch
}
func (p *kubeProxyClientStreams) resize(size *remotecommand.TerminalSize) error {
escape := fmt.Sprintf("\x1b[8;%d;%dt", size.Height, size.Width)
_, err := p.stdout.Write([]byte(escape))
return trace.Wrap(err)
}
func (p *kubeProxyClientStreams) forceTerminate() <-chan struct{} {
return make(chan struct{})
}
func (p *kubeProxyClientStreams) sendStatus(err error) error {
return trace.Wrap(p.proxy.sendStatus(err))
}
func (p *kubeProxyClientStreams) Close() error {
if p.sizeQueue != nil {
p.sizeQueue.Close()
}
p.wg.Wait()
return nil
}
// terminalResizeMessage is a message that contains the terminal size and the source of the resize event.
type terminalResizeMessage struct {
size *remotecommand.TerminalSize
source uuid.UUID
}
// multiResizeQueue is a merged queue of multiple terminal size queues.
type multiResizeQueue struct {
resizes chan terminalResizeMessage
cancels map[string]context.CancelFunc
callback func(terminalResizeMessage)
mutex sync.Mutex
parentCtx context.Context
lastSize *remotecommand.TerminalSize
}
func newMultiResizeQueue(parentCtx context.Context) *multiResizeQueue {
return &multiResizeQueue{
resizes: make(chan terminalResizeMessage),
cancels: make(map[string]context.CancelFunc),
parentCtx: parentCtx,
}
}
func (r *multiResizeQueue) getLastSize() *remotecommand.TerminalSize {
r.mutex.Lock()
defer r.mutex.Unlock()
return r.lastSize
}
// close stops every forwarder. Canceling the parent context has the same effect; this is the explicit teardown path.
func (r *multiResizeQueue) close() {
r.mutex.Lock()
defer r.mutex.Unlock()
for id, cancel := range r.cancels {
cancel()
delete(r.cancels, id)
}
}
func (r *multiResizeQueue) add(id string, queue <-chan terminalResizeMessage) {
r.mutex.Lock()
defer r.mutex.Unlock()
ctx, cancel := context.WithCancel(r.parentCtx)
r.cancels[id] = cancel
go func() {
defer func() {
r.mutex.Lock()
delete(r.cancels, id)
r.mutex.Unlock()
}()
forwardResizes(ctx, queue, r.resizes)
}()
}
func (r *multiResizeQueue) remove(id string) {
r.mutex.Lock()
defer r.mutex.Unlock()
if cancel, ok := r.cancels[id]; ok {
cancel()
delete(r.cancels, id)
}
}
// forwardResizes drains queue into out until the queue is closed or ctx is canceled (the party is removed or the session ends).
func forwardResizes(ctx context.Context, queue <-chan terminalResizeMessage, out chan<- terminalResizeMessage) {
for {
select {
case <-ctx.Done():
return
case msg, ok := <-queue:
if !ok {
return
}
select {
case out <- msg:
case <-ctx.Done():
return
}
}
}
}
func (r *multiResizeQueue) Next() *remotecommand.TerminalSize {
select {
// If the parent context is canceled, the session has ended and we return early.
case <-r.parentCtx.Done():
return nil
case msg := <-r.resizes:
r.callback(msg)
r.mutex.Lock()
r.lastSize = msg.size
r.mutex.Unlock()
return msg.size
}
}
// party represents one participant of the session and their associated state.
type party struct {
Ctx authContext
ID uuid.UUID
Client remoteClient
Mode types.SessionParticipantMode
closeC chan error
closeOnce sync.Once
}
// newParty creates a new party.
func newParty(ctx authContext, mode types.SessionParticipantMode, client remoteClient) *party {
return &party{
Ctx: ctx,
ID: uuid.New(),
Client: client,
Mode: mode,
closeC: make(chan error, 1),
}
}
// InformClose informs the party that he must leave the session.
func (p *party) InformClose(err error) {
p.closeOnce.Do(func() {
p.closeC <- err
close(p.closeC)
})
}
// CloseConnection closes the party underlying connection.
func (p *party) CloseConnection() error {
return trace.Wrap(p.Client.Close())
}
// session represents an ongoing k8s session.
type session struct {
mu sync.RWMutex
// ctx is the auth context of the session initiator
ctx authContext
forwarder *Forwarder
req *http.Request
params httprouter.Params
id uuid.UUID
// parties is a map of currently active parties.
parties map[uuid.UUID]*party
// partiesHistorical is a map of all current previous parties.
// This is used for audit trails.
partiesHistorical map[uuid.UUID]*party
log *slog.Logger
io *srv.TermManager
terminalSizeQueue *multiResizeQueue
tracker *srv.SessionTracker
accessEvaluator moderation.SessionAccessEvaluator
recorder events.SessionPreparerRecorder
emitter apievents.Emitter
podName string
podNamespace string
container string
started bool
initiator uuid.UUID
expires time.Time
// sess is the clusterSession used to establish this session.
sess *clusterSession
closeC chan struct{}
closeOnce sync.Once
// PresenceEnabled is set to true if MFA based presence is required.
PresenceEnabled bool
// Set if we should broadcast information about participant requirements to the session.
displayParticipantRequirements bool
// invitedUsers is a list of users that were invited to the session.
invitedUsers []string
// reason is the reason for the session.
reason string
// weakEventsWaiter is used to wait for events to be emitted and goroutines closed
// when a session is closed.
// Note: this is a weakWaitGroup and doesn't have the same guarantees as sync.WaitGroup.
// Please see the documentation for [weakWaitGroup] for more information.
weakEventsWaiter weakWaitGroup
streamContext context.Context
streamContextCancel context.CancelFunc
// partiesWg is a sync.WaitGroup that tracks the number of active parties
// in this session. It's incremented when a party joins a session and
// decremented when he leaves - it waits until the session leave events
// are emitted for every party before returning.
partiesWg sync.WaitGroup
// terminationErr is set when the session is terminated.
terminationErr error
}
// newSession creates a new session in pending mode.
func newSession(ctx authContext, forwarder *Forwarder, req *http.Request, params httprouter.Params, initiator *party, sess *clusterSession) (*session, error) {
id := uuid.New()
log := forwarder.log.With("session", id.String())
log.DebugContext(req.Context(), "Creating session")
var policySets []*types.SessionTrackerPolicySet
unscopedCtx, isUnscoped := ctx.UnscopedContext()
// TODO(eriktate/scopes): scoped identities don't support policy sets, so we skip attempting to aggregate
// them unless the identity is unscoped.
if isUnscoped {
roles := unscopedCtx.Checker.Roles()
for _, role := range roles {
policySet := role.GetSessionPolicySet()
policySets = append(policySets, &policySet)
}
}
q := req.URL.Query()
accessEvaluator := moderation.NewSessionAccessEvaluator(policySets, types.KubernetesSessionKind, ctx.User.GetName())
if accessEvaluator.IsModerated() && forwarder.cfg.GetScope() != "" {
// If the kube forwarder is scoped then moderated sessions are not supported and access to
// KindKubernetesWaitingContainer will be denied. We need to return an explicit error for unscoped,
// moderated sessions in order to prevent any sort of bypass interacting with kube waiting containers.
return nil, trace.AccessDenied("scoped kubernetes clusters do not support moderated sessions")
}
io := srv.NewTermManager()
streamContext, streamContextCancel := context.WithCancel(forwarder.ctx)
namespace := params.ByName("podNamespace")
podName := params.ByName("podName")
container := q.Get("container")
recorder, err := recorder.New(recorder.Config{
SessionID: tsession.ID(id.String()),
ServerID: forwarder.cfg.HostID,
Namespace: forwarder.cfg.Namespace,
Clock: forwarder.cfg.Clock,
ClusterName: forwarder.cfg.ClusterName,
RecordingCfg: ctx.recordingConfig,
SyncStreamer: forwarder.cfg.AuthClient,
DataDir: forwarder.cfg.DataDir,
Component: teleport.Component(teleport.ComponentSession, teleport.ComponentProxyKube),
// Session stream is using server context, not session context,
// to make sure that session is uploaded even after it is closed
Context: forwarder.ctx,
})
if err != nil {
streamContextCancel()
return nil, trace.Wrap(err)
}
s := &session{
podName: podName,
podNamespace: namespace,
container: container,
ctx: ctx,
forwarder: forwarder,
req: req,
params: params,
id: id,
parties: make(map[uuid.UUID]*party),
partiesHistorical: make(map[uuid.UUID]*party),
log: log,
io: io,
accessEvaluator: accessEvaluator,
terminalSizeQueue: newMultiResizeQueue(streamContext),
started: false,
sess: sess,
closeC: make(chan struct{}),
initiator: initiator.ID,
expires: time.Now().UTC().Add(sessionMaxLifetime),
PresenceEnabled: ctx.Identity.GetIdentity().MFAVerified != "",
displayParticipantRequirements: utils.AsBool(q.Get(teleport.KubeSessionDisplayParticipantRequirementsQueryParam)),
invitedUsers: strings.Split(q.Get(teleport.KubeSessionInvitedQueryParam), ","),
reason: q.Get(teleport.KubeSessionReasonQueryParam),
streamContext: streamContext,
streamContextCancel: streamContextCancel,
partiesWg: sync.WaitGroup{},
// if session ever starts, emitter and recorder will be replaced
// by actual emitter and recorder.
emitter: forwarder.cfg.Emitter,
recorder: recorder,
}
s.io.AddWriter(sessionRecorderID, recorder)
s.io.OnWriteError = s.disconnectPartyOnErr
s.io.OnReadError = s.disconnectPartyOnErr
s.BroadcastMessage("Creating session with ID: %v...", id.String())
go func() {
if _, open := <-s.io.TerminateNotifier(); open {
err := s.Close()
if err != nil {
s.log.ErrorContext(req.Context(), "Failed to close session", "error", err)
}
}
}()
if err := s.trackSession(initiator, policySets); err != nil {
return nil, trace.Wrap(err)
}
return s, nil
}
// disconnectPartyOnErr is called when any party connection returns an error.
// It is used to properly handle client disconnections.
func (s *session) disconnectPartyOnErr(idString string, err error) {
if idString == sessionRecorderID {
s.log.ErrorContext(s.sess.sessionCtx, "Failed to write to session recorder, closing session")
s.Close()
return
}
id, uuidParseErr := uuid.Parse(idString)
if uuidParseErr != nil {
s.log.ErrorContext(s.sess.sessionCtx, "Unable to decode party id",
"party_id", idString,
"error", uuidParseErr,
)
return
}
wasActive, leaveErr := s.leave(s.streamContext, id)
if leaveErr != nil {
s.log.ErrorContext(s.sess.sessionCtx, "Failed to disconnect party from the session",
"party_id", idString,
"error", leaveErr,
)
}
if wasActive {
// log the error only if it was the reason for the user disconnection.
s.log.ErrorContext(s.sess.sessionCtx, "Encountered error with party, disconnecting them from the session",
"error", err,
"party_id", idString,
)
}
}
// checkPresence checks the presence timestamp of involved moderators
// and kicks them if they are not active.
func (s *session) checkPresence(ctx context.Context) error {
s.mu.Lock()
defer s.mu.Unlock()
for _, participant := range s.tracker.GetParticipants() {
if participant.ID == s.initiator.String() {
continue
}
if participant.Mode == string(types.SessionModeratorMode) && time.Now().UTC().After(participant.LastActive.Add(PresenceMaxDifference)) {
s.log.DebugContext(s.sess.sessionCtx, "Participant is not active, kicking", "participant_id", participant.ID)
id, _ := uuid.Parse(participant.ID)
_, err := s.unlockedLeave(ctx, id)
if err != nil {
s.log.WarnContext(s.sess.sessionCtx, "Failed to kick participant for inactivity",
"participant_id", participant.ID,
"error", err,
)
}
}
}
return nil
}
// launch waits until the session meets access requirements and then transitions the session
// to a running state.
func (s *session) launch(ephemeralContainerStatus *corev1.ContainerStatus) (returnErr error) {
defer func() {
err := s.Close()
if err != nil {
s.log.ErrorContext(s.req.Context(), "Failed to close session",
"session_id", s.id,
"error", err,
)
}
}()
s.log.DebugContext(s.req.Context(), "Launching session", "session_id", s.id)
q := s.req.URL.Query()
namespace := s.params.ByName("podNamespace")
podName := s.params.ByName("podName")
container := q.Get("container")
request := &remoteCommandRequest{
podNamespace: namespace,
podName: podName,
containerName: container,
cmd: q["command"],
stdin: utils.AsBool(q.Get("stdin")),
stdout: utils.AsBool(q.Get("stdout")),
stderr: utils.AsBool(q.Get("stderr")),
httpRequest: s.req,
httpResponseWriter: nil,
context: s.req.Context(),
pingPeriod: s.forwarder.cfg.ConnPingPeriod,
}
s.podName = request.podName
s.BroadcastMessage("Connecting to %v over K8S", s.podName)
eventPodMeta := request.eventPodMeta(request.context, s.sess.kubeAPICreds)
onFinished, err := s.lockedSetupLaunch(request, eventPodMeta)
defer func() {
if returnErr != nil {
s.setTerminationErr(returnErr)
s.reportErrorToSessionRecorder(returnErr)
s.log.WarnContext(s.req.Context(), "Executor failed while streaming", "error", returnErr)
}
// call onFinished to emit the session.end and exec events.
// onFinished is never nil.
onFinished(returnErr)
}()
if err != nil {
return trace.Wrap(err)
}
termParams := tsession.TerminalParams{
W: 100,
H: 100,
}
sessionStartEvent, err := s.recorder.PrepareSessionEvent(&apievents.SessionStart{
Metadata: apievents.Metadata{
Type: events.SessionStartEvent,
Code: events.SessionStartCode,
ClusterName: s.forwarder.cfg.ClusterName,
},
ServerMetadata: s.sess.getServerMetadata(),
SessionMetadata: s.getSessionMetadata(),
UserMetadata: s.ctx.eventUserMeta(),
ConnectionMetadata: apievents.ConnectionMetadata{
RemoteAddr: s.req.RemoteAddr,
LocalAddr: s.sess.kubeAddress,
Protocol: events.EventProtocolKube,
},
TerminalSize: termParams.Serialize(),
KubernetesClusterMetadata: s.ctx.eventClusterMeta(s.req),
KubernetesPodMetadata: eventPodMeta,
InitialCommand: q["command"],
SessionRecording: s.ctx.recordingConfig.GetMode(),
Invited: s.invitedUsers,
Reason: s.reason,
})
if err == nil {
if err := s.recorder.RecordEvent(s.forwarder.ctx, sessionStartEvent); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to record session start event", "error", err)
}
if err := s.emitter.EmitAuditEvent(s.forwarder.ctx, sessionStartEvent.GetAuditEvent()); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to emit session start event", "error", err)
}
} else {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to set up session start event - event will not be recorded", "error", err)
}
s.weakEventsWaiter.Add(1)
go func() {
defer s.weakEventsWaiter.Done()
t := time.NewTimer(time.Until(s.expires))
defer t.Stop()
select {
case <-t.C:
s.BroadcastMessage("Session expired, closing...")
err := s.Close()
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to close session", "error", err)
}
case <-s.closeC:
}
}()
if err = s.tracker.UpdateState(s.forwarder.ctx, types.SessionState_SessionStateRunning); err != nil {
s.log.WarnContext(s.forwarder.ctx, "Failed to set tracker state to running", "error", err)
}
executor, executorCleanup, err := s.forwarder.getExecutor(s.sess, s.req)
if err != nil {
s.log.WarnContext(s.forwarder.ctx, "Failed creating executor", "error", err)
return trace.Wrap(err)
}
defer executorCleanup()
options := remotecommand.StreamOptions{
Stdin: s.io,
Stdout: s.io,
Stderr: s.io,
Tty: true,
TerminalSizeQueue: s.terminalSizeQueue,
}
s.io.On()
// If the container is ephemeral and already terminated, we should
// retrieve the logs and return early.
if ephemeralContainerStatus != nil && ephemeralContainerStatus.State.Terminated != nil {
err := s.retrieveAlreadyStoppedPodLogs(
namespace,
podName,
container,
)
return trace.Wrap(err)
}
if streamErr := executor.StreamWithContext(s.streamContext, options); streamErr != nil {
// If the container isn't ephemeral, return the error.
if ephemeralContainerStatus == nil {
return trace.Wrap(streamErr)
}
fmt.Fprintf(s.io, "\r\nwarning: couldn't attach to pod/%s, falling back to streaming logs: %v\r\n", podName, streamErr)
err := s.retrieveAlreadyStoppedPodLogs(
namespace,
podName,
container,
)
return trace.Wrap(err)
}
return nil
}
func (s *session) setTerminationErr(err error) {
s.mu.Lock()
defer s.mu.Unlock()
s.setTerminationErrUnlocked(err)
}
func (s *session) setTerminationErrUnlocked(err error) {
if s.terminationErr != nil {
return
}
s.terminationErr = err
}
// reportErrorToSessionRecorder reports the error to the session recorder
// if it is set.
func (s *session) reportErrorToSessionRecorder(err error) {
if err == nil {
return
}
if s.recorder != nil {
fmt.Fprintf(s.recorder, "\r\n---\r\nSession exited with error: %v\r\n", err)
}
}
func (s *session) lockedSetupLaunch(request *remoteCommandRequest, eventPodMeta apievents.KubernetesPodMetadata) (func(error), error) {
s.mu.Lock()
defer s.mu.Unlock()
s.started = true
sessionStart := s.forwarder.cfg.Clock.Now().UTC()
if s.sess.isLocalKubernetesCluster {
s.terminalSizeQueue.callback = func(termSize terminalResizeMessage) {
s.mu.Lock()
defer s.mu.Unlock()
for id, p := range s.parties {
// Skip the party that sent the resize event to avoid a resize loop.
if p.Client.queueID() == termSize.source {
continue
}
err := p.Client.resize(termSize.size)
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to resize participant",
"party_id", id.String(),
"error", err,
)
}
}
params := tsession.TerminalParams{
W: int(termSize.size.Width),
H: int(termSize.size.Height),
}
resizeEvent, err := s.recorder.PrepareSessionEvent(&apievents.Resize{
Metadata: apievents.Metadata{
Type: events.ResizeEvent,
Code: events.TerminalResizeCode,
ClusterName: s.forwarder.cfg.ClusterName,
},
ConnectionMetadata: apievents.ConnectionMetadata{
RemoteAddr: s.req.RemoteAddr,
Protocol: events.EventProtocolKube,
},
ServerMetadata: s.sess.getServerMetadata(),
SessionMetadata: s.getSessionMetadata(),
UserMetadata: s.ctx.eventUserMeta(),
TerminalSize: params.Serialize(),
KubernetesClusterMetadata: s.ctx.eventClusterMeta(s.req),
KubernetesPodMetadata: eventPodMeta,
})
if err == nil {
// Report the updated window size to the event log (this is so the sessions
// can be replayed correctly).
if err := s.recorder.RecordEvent(s.forwarder.ctx, resizeEvent); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to emit terminal resize event", "error", err)
}
} else {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to set up terminal resize event - event will not be recorded", "error", err)
}
}
} else {
s.terminalSizeQueue.callback = func(resize terminalResizeMessage) {}
}
// If we get here, it means we are going to have a session.end event.
// This increments the waiter so that session.Close() guarantees that once called
// the events are emitted before closing the emitter/recorder.
// It might happen when a user disconnects or when a moderator forces an early
// termination.
s.weakEventsWaiter.Add(1)
onFinish := func(errExec error) {
defer s.weakEventsWaiter.Done()
s.mu.Lock()
defer s.mu.Unlock()
serverMetadata := s.sess.getServerMetadata()
sessionMetadata := s.getSessionMetadata()
conMetadata := apievents.ConnectionMetadata{
RemoteAddr: s.req.RemoteAddr,
LocalAddr: s.sess.kubeAddress,
Protocol: events.EventProtocolKube,
}
execEvent := &apievents.Exec{
Metadata: apievents.Metadata{
Type: events.ExecEvent,
ClusterName: s.forwarder.cfg.ClusterName,
// can be changed to ExecFailureCode if errExec is not nil
Code: events.ExecCode,
},
ServerMetadata: serverMetadata,
SessionMetadata: sessionMetadata,
UserMetadata: s.sess.eventUserMeta(),
ConnectionMetadata: conMetadata,
CommandMetadata: apievents.CommandMetadata{
Command: strings.Join(request.cmd, " "),
},
KubernetesClusterMetadata: s.ctx.eventClusterMeta(s.req),
KubernetesPodMetadata: eventPodMeta,
}
if errExec != nil {
execEvent.Code = events.ExecFailureCode
execEvent.Error, execEvent.ExitCode = exitCode(errExec)
}
if err := s.emitter.EmitAuditEvent(s.forwarder.ctx, execEvent); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to emit exec event", "error", err)
}
sessionDataEvent := &apievents.SessionData{
Metadata: apievents.Metadata{
Type: events.SessionDataEvent,
Code: events.SessionDataCode,
ClusterName: s.forwarder.cfg.ClusterName,
},
ServerMetadata: serverMetadata,
SessionMetadata: sessionMetadata,
UserMetadata: s.sess.eventUserMeta(),
ConnectionMetadata: conMetadata,
// Bytes transmitted from user to pod.
BytesTransmitted: s.io.CountRead(),
// Bytes received from pod by user.
BytesReceived: s.io.CountWritten(),
}
if err := s.emitter.EmitAuditEvent(s.forwarder.ctx, sessionDataEvent); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to emit session data event", "error", err)
}
sessionEndEvent, err := s.recorder.PrepareSessionEvent(&apievents.SessionEnd{
Metadata: apievents.Metadata{
Type: events.SessionEndEvent,
Code: events.SessionEndCode,
ClusterName: s.forwarder.cfg.ClusterName,
},
ServerMetadata: serverMetadata,
SessionMetadata: sessionMetadata,
UserMetadata: s.sess.eventUserMeta(),
ConnectionMetadata: conMetadata,
Interactive: true,
Participants: s.allParticipants(),
StartTime: sessionStart,
EndTime: s.forwarder.cfg.Clock.Now().UTC(),
KubernetesClusterMetadata: s.ctx.eventClusterMeta(s.req),
KubernetesPodMetadata: eventPodMeta,
InitialCommand: request.cmd,
SessionRecording: s.ctx.recordingConfig.GetMode(),
})
if err == nil {
if err := s.recorder.RecordEvent(s.forwarder.ctx, sessionEndEvent); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to record session end event", "error", err)
}
if err := s.emitter.EmitAuditEvent(s.forwarder.ctx, sessionEndEvent.GetAuditEvent()); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to emit session end event", "error", err)
}
} else {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to set up session end event - event will not be recorded", "error", err)
}
}
// If the identity is verified with an MFA device, we enabled MFA-based presence for the session.
if s.PresenceEnabled {
s.weakEventsWaiter.Add(1)
go func() {
defer s.weakEventsWaiter.Done()
ticker := time.NewTicker(PresenceVerifyInterval)
defer ticker.Stop()
for {
select {
case <-ticker.C:
err := s.checkPresence(s.streamContext)
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to check presence, closing session as a security measure", "error", err)
if err := s.Close(); err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to close session", "error", err)
}
return
}
case <-s.closeC:
return
}
}
}()
}
return onFinish, nil
}
// join attempts to connect a party to the session.
func (s *session) join(ctx context.Context, p *party, emitJoinEvent bool) error {
if p.Ctx.User.GetName() != s.ctx.User.GetName() {
unscopedCtx, isUnscoped := p.Ctx.UnscopedContext()
if !isUnscoped {
return trace.Wrap(services.ErrScopedIdentity, "joining moderated session")
}
roles := unscopedCtx.Checker.Roles()
accessContext := moderation.SessionAccessContext{
Username: p.Ctx.User.GetName(),
Roles: roles,
}
modes := s.accessEvaluator.CanJoin(accessContext)
if !slices.Contains(modes, p.Mode) {
return trace.AccessDenied("insufficient permissions to join session")
}
}
if s.tracker.GetState() == types.SessionState_SessionStateTerminated {
return trace.AccessDenied("The requested session is not active")
}
s.log.DebugContext(s.forwarder.ctx, "Tracking participant", "participant_id", p.ID)
participant := &types.Participant{
ID: p.ID.String(),
User: p.Ctx.Identity.GetIdentity().Username,
Cluster: p.Ctx.Identity.GetIdentity().OriginClusterName,
Mode: string(p.Mode),
LastActive: time.Now().UTC(),
}
if err := s.tracker.AddParticipant(s.forwarder.ctx, participant); err != nil {
return trace.Wrap(err)
}
// We only want to emit the session.join when someone tries to join a session via
// tsh kube join and not when the original session owner terminal streams are
// connected to the Kubernetes session.
if emitJoinEvent {
s.emitSessionJoinEvent(p)
}
recentWrites := s.io.GetRecentHistory()
if _, err := p.Client.stdoutStream().Write(recentWrites); err != nil {
s.log.WarnContext(s.forwarder.ctx, "Failed to write history to participant", "error", err)
}
s.BroadcastMessage("User %v joined the session with participant mode: %v.", p.Ctx.User.GetName(), p.Mode)
// increment the party track waitgroup.
// It is decremented when session.leave() finishes its execution.
s.partiesWg.Add(1)
s.mu.Lock()
defer s.mu.Unlock()
stringID := p.ID.String()
s.parties[p.ID] = p
s.partiesHistorical[p.ID] = p
s.terminalSizeQueue.add(stringID, p.Client.resizeQueue())
// If the session is already running, we need to resize the new party's terminal
// to match the last terminal size.
// This is done to ensure that the new party's terminal is the same size as the
// other parties' terminals and no discrepancies are present.
if lastQueueSize := s.terminalSizeQueue.getLastSize(); lastQueueSize != nil {
if err := p.Client.resize(lastQueueSize); err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to resize participant",
"participant_id", stringID,
"error", err,
)
}
}
if p.Mode == types.SessionPeerMode {
s.io.AddReader(stringID, p.Client.stdinStream())
}
s.io.AddWriter(stringID, p.Client.stdoutStream())
// Send the participant mode and controls to the additional participant
if p.Ctx.User.GetName() != s.ctx.User.GetName() {
err := srv.MsgParticipantCtrls(p.Client.stdoutStream(), p.Mode)
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Could not send intro message to participant",
"error", err,
"participant_id", stringID,
)
}
}
// Allow the moderator to force terminate the session
if p.Mode == types.SessionModeratorMode {
s.weakEventsWaiter.Add(1)
go func() {
defer s.weakEventsWaiter.Done()
c := p.Client.forceTerminate()
select {
case <-c:
s.setTerminationErr(sessionTerminatedByModeratorErr)
go func() {
s.log.DebugContext(s.forwarder.ctx, "Received force termination request")
err := s.Close()
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to close session", "error", err)
}
}()
case <-s.closeC:
return
}
}()
}
// Detach cancellation: the participant is already registered above, so a
// canceled request context here would leak them in s.parties/partiesWg
// because the caller treats this error as fatal and never calls leave.
canStart, _, err := s.canStart(context.WithoutCancel(ctx))
if err != nil {
return trace.Wrap(err)
}
if !s.started {
if canStart {
// create an ephemeral container if this session will be
// running in one now that the moderated session is approved
startedEphemeralCont, err := s.createEphemeralContainer()
if err != nil {
// if the ephemeral container creation fails, close the session
// and return the error. We need to close the session here because
// we must inform all parties that the session is closing.
s.setTerminationErrUnlocked(err)
s.reportErrorToSessionRecorder(err)
s.log.WarnContext(s.forwarder.ctx, "Executor failed while creating ephemeral pod", "error", err)
go func() {
err := s.Close()
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to close session", "error", err)
}
}()
return trace.Wrap(err)
}
go func() {
if err := s.launch(startedEphemeralCont); err != nil {
s.log.WarnContext(s.forwarder.ctx, "Failed to launch Kubernetes session", "error", err)
}
}()
} else if len(s.parties) == 1 {
const base = "Waiting for required participants..."
if s.displayParticipantRequirements {
s.BroadcastMessage(base+"\r\n%v", s.accessEvaluator.PrettyRequirementsList())
} else {
s.BroadcastMessage(base)
}
}
} else if canStart && s.tracker.GetState() == types.SessionState_SessionStatePending {
// If the session is already running, but the party is a moderator that left
// a session with onLeave=pause and then rejoined, we need to unpause the session.
// When the moderator left the session, the session was paused, and we spawn
// a goroutine to wait for the moderator to rejoin. If the moderator rejoins
// before the session ends, we need to unpause the session by updating its state and
// the goroutine will unblock the s.io terminal.
// types.SessionState_SessionStatePending marks a session that is waiting for
// a moderator to rejoin.
if err := s.tracker.UpdateState(s.forwarder.ctx, types.SessionState_SessionStateRunning); err != nil {
s.log.WarnContext(s.forwarder.ctx, "Failed to update tracker to running state")
}
}
return nil
}
// createEphemeralContainer creates an ephemeral container and waits for it to start.
func (s *session) createEphemeralContainer() (*corev1.ContainerStatus, error) {
initUser := s.parties[s.initiator]
username := initUser.Ctx.Identity.GetIdentity().Username
namespace := s.params.ByName("podNamespace")
podName := s.params.ByName("podName")
container := s.req.URL.Query().Get("container")
if s.forwarder.cfg.GetScope() != "" {
// If the kube forwarder is scoped then moderated sessions are not supported and access to
// KindKubernetesWaitingContainer will be denied. We need to return without error to prevent
// interactive exec from failing
return nil, nil
}
waitingCont, err := s.forwarder.cfg.CachingAuthClient.GetKubernetesWaitingContainer(
s.forwarder.ctx,
kubewaitingcontainerpb.GetKubernetesWaitingContainerRequest_builder{
Username: username,
Cluster: s.ctx.kubeClusterName,
Namespace: namespace,
PodName: podName,
ContainerName: container,
}.Build(),
)
if trace.IsNotFound(err) {
return nil, nil
} else if err != nil {
return nil, trace.Wrap(err)
}
if err = s.forwarder.cfg.AuthClient.DeleteKubernetesWaitingContainer(
s.forwarder.ctx,
kubewaitingcontainerpb.DeleteKubernetesWaitingContainerRequest_builder{
Username: username,
Cluster: s.ctx.kubeClusterName,
Namespace: namespace,
PodName: podName,
ContainerName: container,
}.Build(),
); err != nil {
return nil, trace.Wrap(err)
}
s.log.DebugContext(s.forwarder.ctx, "Creating ephemeral container on pod", "container", container, "pod", podName)
containerStatus, err := s.patchAndWaitForPodEphemeralContainer(s.forwarder.ctx, &initUser.Ctx, s.req.Header, waitingCont)
return containerStatus, trace.Wrap(err)
}
func (s *session) BroadcastMessage(format string, args ...any) {
if s.accessEvaluator.IsModerated() {
s.io.BroadcastMessage(fmt.Sprintf(format, args...))
}
}
// emitSessionJoinEvent emits a session.join audit event when a user joins
// the session.
// This function requires that the session must be active, otherwise audit logger
// will discard the event.
func (s *session) emitSessionJoinEvent(p *party) {
sessionJoinEvent := &apievents.SessionJoin{
Metadata: apievents.Metadata{
Type: events.SessionJoinEvent,
Code: events.SessionJoinCode,
ClusterName: s.ctx.teleportCluster.name,
},
KubernetesClusterMetadata: apievents.KubernetesClusterMetadata{
KubernetesCluster: s.ctx.kubeClusterName,
// joining moderators, obervers and peers don't have any
// kubernetes metadata configured.
KubernetesUsers: []string{},
KubernetesGroups: []string{},
KubernetesLabels: s.ctx.kubeClusterLabels,
},
SessionMetadata: s.getSessionMetadata(),
UserMetadata: p.Ctx.eventUserMetaWithLogin("root"),
ConnectionMetadata: apievents.ConnectionMetadata{
RemoteAddr: s.params.ByName("podName"),
},
}
s.prepareAndEmitEvent(sessionJoinEvent)
}
func (s *session) prepareAndEmitEvent(evt apievents.AuditEvent) {
preparedEvent, err := s.recorder.PrepareSessionEvent(evt)
if err == nil {
evt = preparedEvent.GetAuditEvent()
if err := s.recorder.RecordEvent(s.forwarder.ctx, preparedEvent); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to record event", "error", err, "event_type", evt.GetType())
}
} else {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to prepare event - event will not be recorded into session recording.", "error", err, "event_type", evt.GetType())
}
// Always emit the event to the audit log, even if preparing
if err := s.emitter.EmitAuditEvent(s.forwarder.ctx, evt); err != nil {
s.forwarder.log.WarnContext(s.forwarder.ctx, "Failed to emit event to Audit Log.", "error", err, "event_type", evt.GetType())
}
}
// leave removes a party from the session and returns if the party was still active
// in the session. If the party wasn't found, it returns false, nil.
func (s *session) leave(ctx context.Context, id uuid.UUID) (bool, error) {
s.mu.Lock()
defer s.mu.Unlock()
return s.unlockedLeave(ctx, id)
}
// unlockedLeave removes a party from the session without locking the mutex.
// The boolean returned identifies if the party was still active in the session.
// If the party wasn't found, it returns false, nil.
// In order to call this function, lock the mutex before.
func (s *session) unlockedLeave(ctx context.Context, id uuid.UUID) (bool, error) {
var errs []error
stringID := id.String()
party := s.parties[id]
if party == nil {
return false, nil
}
// Waits until the function execution ends to release the parties waitgroup.
// It's used to prevent the session to terminate the events emitter before
// the session leave event is emitted.
defer s.partiesWg.Done()
delete(s.parties, id)
s.terminalSizeQueue.remove(stringID)
s.io.DeleteReader(stringID)
s.io.DeleteWriter(stringID)
s.BroadcastMessage("User %v left the session.", party.Ctx.User.GetName())
sessionLeaveEvent := &apievents.SessionLeave{
Metadata: apievents.Metadata{
Type: events.SessionLeaveEvent,
Code: events.SessionLeaveCode,
ClusterName: s.ctx.teleportCluster.name,
},
SessionMetadata: s.getSessionMetadata(),
UserMetadata: party.Ctx.eventUserMetaWithLogin("root"),
ConnectionMetadata: apievents.ConnectionMetadata{
RemoteAddr: s.params.ByName("podName"),
},
}
s.prepareAndEmitEvent(sessionLeaveEvent)
s.log.DebugContext(s.forwarder.ctx, "No longer tracking participant", "participant_id", party.ID)
err := s.tracker.RemoveParticipant(s.forwarder.ctx, party.ID.String())
if err != nil {
errs = append(errs, trace.Wrap(err))
}
party.InformClose(s.terminationErr)
if len(s.parties) == 0 || id == s.initiator {
go func() {
// Currently, Teleport closes the session when the initiator exits.
// So, it is safe to remove it
s.forwarder.deleteSession(s.id)
// close session
err := s.Close()
if err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to close session", "error", err)
}
}()
return true, trace.NewAggregate(errs...)
}
// We wait until here to return to check if we should terminate the
// session.
if len(errs) > 0 {
return true, trace.NewAggregate(errs...)
}
canStart, options, err := s.canStart(ctx)
if err != nil {
return true, trace.Wrap(err)
}
if !canStart {
if options.OnLeaveAction == types.OnSessionLeaveTerminate {
go func() {
if err := s.Close(); err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to close session", "error", err)
}
}()
return true, nil
}
// pause session and wait for another party to resume
s.io.Off()
s.BroadcastMessage("Session paused, Waiting for required participants...")
if err := s.tracker.UpdateState(s.forwarder.ctx, types.SessionState_SessionStatePending); err != nil {
s.log.WarnContext(s.forwarder.ctx, "Failed to set tracker state to pending")
}
go func() {
if state := s.tracker.WaitForStateUpdate(types.SessionState_SessionStatePending); state == types.SessionState_SessionStateRunning {
s.BroadcastMessage("Resuming session...")
s.io.On()
}
}()
}
return true, nil
}
// allParticipants returns a list of all historical participants of the session.
func (s *session) allParticipants() []string {
var participants []string
for _, p := range s.partiesHistorical {
username := services.UsernameForCluster(
services.UsernameForClusterConfig{
User: p.Ctx.Identity.GetIdentity().Username,
OriginClusterName: p.Ctx.Identity.GetIdentity().OriginClusterName,
LocalClusterName: p.Ctx.Identity.GetIdentity().TeleportCluster,
},
)
participants = append(participants, username)
}
return participants
}
// canStart checks if a session can start with the current set of participants.
func (s *session) canStart(ctx context.Context) (bool, moderation.PolicyOptions, error) {
var participants []moderation.SessionAccessContext
for _, party := range s.parties {
if party.Ctx.User.GetName() == s.ctx.User.GetName() {
continue
}
roleNames := party.Ctx.Identity.GetIdentity().Groups
roles, err := getRolesByName(ctx, s.forwarder, roleNames)
if err != nil {
return false, moderation.PolicyOptions{}, trace.Wrap(err)
}
participants = append(participants, moderation.SessionAccessContext{
Username: party.Ctx.User.GetName(),
Roles: roles,
Mode: party.Mode,
})
}
yes, options, err := s.accessEvaluator.FulfilledFor(participants)
return yes, options, trace.Wrap(err)
}
// Close terminates a session and disconnects all participants.
func (s *session) Close() error {
s.closeOnce.Do(func() {
s.BroadcastMessage("Closing session...")
s.io.Close()
// Once tracker is closed parties cannot join the session.
// check session.join for logic.
if err := s.tracker.Close(s.forwarder.ctx); err != nil {
s.log.DebugContext(s.forwarder.ctx, "Failed to close session tracker", "error", err)
}
s.mu.Lock()
terminationErr := s.terminationErr
// terminate all active parties in the session.
for _, party := range s.parties {
party.InformClose(terminationErr)
}
recorder := s.recorder
s.mu.Unlock()
s.log.DebugContext(s.forwarder.ctx, "Closing session", "session_id", logutils.StringerAttr(s.id))
close(s.closeC)
// Wait until every party leaves the session and emits the session leave
// event before closing the recorder - if available.
s.partiesWg.Wait()
s.streamContextCancel()
s.terminalSizeQueue.close()
if recorder != nil {
// wait for events to be emitted before closing the recorder/emitter.
// If we close it immediately we will lose session.end events.
s.weakEventsWaiter.Wait()
if err := recorder.Complete(s.forwarder.ctx); err != nil {
s.log.ErrorContext(s.forwarder.ctx, "Failed to complete session recorder", "error", err)
}
}
})
return nil
}
func getRolesByName(ctx context.Context, forwarder *Forwarder, roleNames []string) ([]types.Role, error) {
var roles []types.Role
for _, roleName := range roleNames {
role, err := forwarder.cfg.CachingAuthClient.GetRole(ctx, roleName)
if err != nil {
return nil, trace.Wrap(err)
}
roles = append(roles, role)
}
return roles, nil
}
// trackSession creates a new session tracker for the kube session.
// While ctx is open, the session tracker's expiration will be extended
// on an interval until the session tracker is closed.
func (s *session) trackSession(p *party, policySet []*types.SessionTrackerPolicySet) error {
ctx := s.req.Context()
command := s.req.URL.Query()["command"]
if len(command) == 0 {
command = s.retrieveEphemeralContainerCommand(ctx, p.Ctx.User.GetName(), s.req.URL.Query().Get("container"))
}
trackerSpec := types.SessionTrackerSpecV1{
SessionID: s.id.String(),
Kind: string(types.KubernetesSessionKind),
State: types.SessionState_SessionStatePending,
Hostname: path.Join(s.podNamespace, s.podName),
ClusterName: s.ctx.teleportCluster.name,
KubernetesCluster: s.ctx.kubeClusterName,
HostUser: p.Ctx.User.GetName(),
HostPolicies: policySet,
Login: "root",
Created: s.forwarder.cfg.Clock.Now(),
Reason: s.reason,
Invited: s.invitedUsers,
HostID: s.forwarder.cfg.HostID,
InitialCommand: command,
}
s.log.DebugContext(ctx, "Creating session tracker")
sessionTrackerService := s.forwarder.cfg.AuthClient
tracker, err := srv.NewSessionTracker(ctx, trackerSpec, sessionTrackerService)
switch {
// there was an error creating the tracker for a moderated session - terminate the session
case err != nil && s.accessEvaluator.IsModerated():
s.log.WarnContext(ctx, "Failed to create session tracker, unable to proceed for moderated session", "error", err)
return trace.Wrap(err)
// there was an error creating the tracker for a non-moderated session - permit the session with a local tracker
case err != nil && !s.accessEvaluator.IsModerated():
s.log.WarnContext(ctx, "Failed to create session tracker, proceeding with local session tracker for non-moderated session")
localTracker, err := srv.NewSessionTracker(ctx, trackerSpec, nil)
// this error means there are problems with the trackerSpec, we need to return it
if err != nil {
return trace.Wrap(err)
}
s.tracker = localTracker
// there was an error even though the tracker wasn't being propagated - return it
case err != nil:
return trace.Wrap(err)
// the tracker was created successfully
default:
s.tracker = tracker
}
go func() {
if err := s.tracker.UpdateExpirationLoop(s.forwarder.ctx, s.forwarder.cfg.Clock); err != nil {
s.log.WarnContext(ctx, "Failed to update session tracker expiration", "error", err)
}
}()
return nil
}
func (s *session) getSessionMetadata() apievents.SessionMetadata {
return s.ctx.Identity.GetIdentity().GetSessionMetadata(s.id.String())
}
// patchPodWithEphemeralContainer creates an ephemeral container and waits
// for it to start.
func (s *session) patchAndWaitForPodEphemeralContainer(
ctx context.Context,
authCtx *authContext,
headers http.Header,
waitingCont *kubewaitingcontainerpb.KubernetesWaitingContainer,
) (containerStatus *corev1.ContainerStatus, err error) {
fmt.Fprintf(s.io, "\r\nCreating ephemeral container %s in pod %s/%s\r\n", waitingCont.GetSpec().GetContainerName(), waitingCont.GetSpec().GetNamespace(), waitingCont.GetSpec().GetPodName())
clientSet, _, err := s.forwarder.impersonatedKubeClient(authCtx, headers)
if err != nil {
return nil, trace.Wrap(err)
}
podClient := clientSet.CoreV1().Pods(authCtx.metaResource.requestedResource.namespace)
result, err := podClient.Patch(ctx,
waitingCont.GetSpec().GetPodName(),
apimachinerytypes.StrategicMergePatchType,
waitingCont.GetSpec().GetPatch(),
metav1.PatchOptions{},
"ephemeralcontainers")
if err != nil {
return nil, trace.Wrap(err)
}
fmt.Fprintf(s.io, "Pod %s/%s successfully patched. Waiting for container to become ready.\r\n",
waitingCont.GetSpec().GetNamespace(),
waitingCont.GetSpec().GetPodName())
fieldSelector := fields.OneTermEqualSelector("metadata.name", waitingCont.GetSpec().GetPodName()).String()
lw := &cache.ListWatch{
ListFunc: func(options metav1.ListOptions) (runtime.Object, error) {
options.FieldSelector = fieldSelector
options.ResourceVersion = result.GetResourceVersion()
options.ResourceVersionMatch = metav1.ResourceVersionMatchNotOlderThan
return podClient.List(ctx, options)
},
WatchFunc: func(options metav1.ListOptions) (watch.Interface, error) {
options.FieldSelector = fieldSelector
options.ResourceVersion = result.GetResourceVersion()
options.ResourceVersionMatch = metav1.ResourceVersionMatchNotOlderThan
return podClient.Watch(ctx, options)
},
}
_, err = watchtools.UntilWithSync(ctx, lw, &corev1.Pod{}, nil, func(ev watch.Event) (bool, error) {
switch ev.Type {
case watch.Deleted:
return false, trace.NotFound("pod %s not found", waitingCont.GetSpec().GetPodName())
}
p, ok := ev.Object.(*corev1.Pod)
if !ok {
return false, trace.BadParameter("watch did not return a pod: %v", ev.Object)
}
s := getEphemeralContainerStatusByName(p, waitingCont.GetSpec().GetContainerName())
if s == nil {
return false, nil
}
if s.State.Running != nil || s.State.Terminated != nil {
containerStatus = s
return true, nil
}
return false, nil
})
if err != nil {
return nil, trace.Wrap(err)
}
fmt.Fprintf(s.io, "Ephemeral container %s is ready.\r\n", waitingCont.GetSpec().GetContainerName())
return containerStatus, nil
}
// retrieveAlreadyStoppedPodLogs retrieves the logs of a stopped pod and writes them to the session's io writer.
func (s *session) retrieveAlreadyStoppedPodLogs(namespace, podName, container string) error {
// If attaching to the container failed, check if the container
// is terminated. If it is, try to stream the logs. If it's not
// terminated or can't be found return the original error.
clientSet, _, err := s.forwarder.impersonatedKubeClient(&s.sess.authContext, s.req.Header)
if err != nil {
return trace.Wrap(err)
}
podClient := clientSet.CoreV1().Pods(namespace)
fmt.Fprintf(s.io, "Failed to attach to the container, attempting to stream logs instead...\r\n")
req := podClient.GetLogs(podName, &corev1.PodLogOptions{Container: container})
r, err := req.Stream(s.streamContext)
if err != nil {
return trace.Wrap(err)
}
if _, err := io.Copy(s.io, r); err != nil {
_ = r.Close()
return trace.Wrap(err)
}
return trace.Wrap(r.Close())
}
// retrieveEphemeralContainerCommand retrieves the command of an ephemeral container
// if it exists.
func (s *session) retrieveEphemeralContainerCommand(ctx context.Context, username, containerName string) []string {
containers, err := s.forwarder.getUserEphemeralContainersForPod(ctx, username, s.ctx.kubeClusterName, s.podNamespace, s.podName)
if err != nil {
s.log.WarnContext(ctx, "Failed to retrieve ephemeral containers", "error", err)
return nil
}
if len(containers) == 0 {
return nil
}
for _, container := range containers {
if container.GetMetadata().GetName() != containerName {
continue
}
contentType, err := patchTypeToContentType(apimachinerytypes.PatchType(container.GetSpec().GetPatchType()))
if err != nil {
return nil
}
encoder, decoder, err := newEncoderAndDecoderForContentType(
contentType,
newClientNegotiator(s.sess.codecFactory),
)
if err != nil {
s.log.WarnContext(ctx, "Failed to create encoder and decoder", "error", err)
return nil
}
currentPod, err := s.forwarder.getPodForEphemeralPatch(
ctx,
&s.ctx,
impersonationHeadersFromWaitingContainer(container),
s.podNamespace,
s.podName,
)
if err != nil {
s.log.WarnContext(ctx, "Failed to get pod for ephemeral patch", "error", err)
return nil
}
pod, _, err := s.forwarder.mergeEphemeralPatchWithCurrentPod(
currentPod,
mergeEphemeralPatchWithCurrentPodConfig{
decoder: decoder,
encoder: encoder,
podPatch: container.GetSpec().GetPatch(),
patchType: apimachinerytypes.PatchType(container.GetSpec().GetPatchType()),
},
)
if err != nil {
s.log.WarnContext(ctx, "Failed to merge ephemeral patch with current pod", "error", err)
return nil
}
for _, ephemeral := range pod.Spec.EphemeralContainers {
if ephemeral.Name == containerName {
return ephemeral.Command
}
}
}
return nil
}
// weakWaitGroup is a specialized synchronization primitive similar to sync.WaitGroup
// but with **relaxed** guarantees. Unlike sync.WaitGroup, weakWaitGroup does not ensure
// that the Wait() method will wait for all Add() calls to reach completion through Done()
// if they are called concurrently. This means that there is a potential leak in the
// synchronization of goroutines that are added to the weakWaitGroup and may be started
// after the Wait() method is called.
//
// Use Case:
// This weakWaitGroup is intended for scenarios where goroutines are initiated from
// various parts of the codebase concurrently and need to be awaited only if they started before
// a certain point in time, specifically before session.Close() is called. If a goroutine
// is initiated after session.Close() has been invoked, it will not be included in the wait process.
// It's the caller responsibility to ensure that all goroutines started after Wait() returns end
// up being a no-op.
//
// Important Considerations:
// - This implementation is UNSAFE as a general-purpose synchronization primitive.
// - It does not guarantee that Wait() will account for all Add() calls, leading to potential
// race conditions or goroutines that may not be properly awaited.
// - Due to these limitations, weakWaitGroup should be used with extreme caution and only
// in contexts where its relaxed guarantees are acceptable and safe.
//
// WARNING:
// This is not a substitute for sync.WaitGroup in situations requiring strong synchronization
// guarantees.
type weakWaitGroup struct {
cond sync.Cond
mu sync.Mutex
count int
}
func (c *weakWaitGroup) Add(delta int) {
c.mu.Lock()
defer c.mu.Unlock()
c.count += delta
}
func (c *weakWaitGroup) Done() {
c.mu.Lock()
defer c.mu.Unlock()
c.count--
if c.count == 0 && c.cond.L != nil {
c.cond.Broadcast()
}
}
func (c *weakWaitGroup) Wait() {
c.mu.Lock()
defer c.mu.Unlock()
if c.count == 0 {
return
}
if c.cond.L == nil {
c.cond.L = &c.mu
}
for c.count > 0 {
c.cond.Wait()
}
}
/*
* Teleport
* Copyright (C) 2025 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 proxy
import (
"encoding/base64"
"net/http"
"strings"
"unicode/utf8"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/tlsca"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
const (
// paramTeleportCluster is the path parameter key containing a base64
// encoded Teleport cluster name for path-routed forwarding.
paramTeleportCluster = "base64Cluster"
// paramKubernetesCluster is the path parameter key containing a base64
// encoded Teleport cluster name for path-routed forwarding.
paramKubernetesCluster = "base64KubeCluster"
)
// parseRouteFromPath extracts route information from the given path parameters
// using constant-defined parameter keys.
func parseRouteFromPath(p httprouter.Params) (string, string, error) {
encodedTeleportCluster := p.ByName(paramTeleportCluster)
if encodedTeleportCluster == "" {
return "", "", trace.BadParameter("no Teleport cluster name found in path")
}
decodedTeleportCluster, err := base64.RawURLEncoding.DecodeString(encodedTeleportCluster)
if err != nil {
return "", "", trace.Wrap(err)
}
encodedKubernetesCluster := p.ByName(paramKubernetesCluster)
if encodedKubernetesCluster == "" {
return "", "", trace.BadParameter("no Kubernetes cluster name found in path")
}
decodedKubernetesCluster, err := base64.RawURLEncoding.DecodeString(encodedKubernetesCluster)
if err != nil {
return "", "", trace.Wrap(err)
}
if !utf8.Valid(decodedTeleportCluster) {
return "", "", trace.BadParameter("invalid Teleport cluster name")
}
if !utf8.Valid(decodedKubernetesCluster) {
return "", "", trace.BadParameter("invalid Kubernetes cluster name")
}
return string(decodedTeleportCluster), string(decodedKubernetesCluster), nil
}
// ensureRouteNotOverwritten checks that the path routing parameters do not
// overwrite any existing RouteToCluster or KubernetesCluster fields in the
// identity, if those fields are set. Additionally, it is valid for the path
// route to be equal to the existing value.
//
// This requirement ensures temporary certs issued for session MFA remain bound
// to the cluster for which they were initially issued and path routing cannot
// be used to access a different target cluster. If MFA is required, path routed
// requests will receive an `ErrSessionMFARequired` as usual and will need to
// request certificates with identity-based routing information. Once the
// temporary identity is issued, the request can proceed as usual through this
// path-based route so long as the path and identity route fields are equal.
func ensureRouteNotOverwritten(ident *tlsca.Identity, routeToCluster, kubernetesCluster string) error {
teleportClusterChanged := ident.RouteToCluster != routeToCluster
kubeClusterChanged := ident.KubernetesCluster != kubernetesCluster
// If session MFA is enabled, either cluster-wide or for the target cluster
// via role options, access attempts without an MFA assertion will pass
// through here and fail during `CheckAccess()` in `authorize()`. If retried
// with an assertion, we should not allow routing parameters to be
// overwritten even if somehow empty ("") as that would allow MFA certs to
// access any cluster.
if ident.MFAVerified != "" && (teleportClusterChanged || kubeClusterChanged) {
return trace.AccessDenied("identity routing parameters are required when MFA assertions are present")
}
const overwriteDeniedMsg = "existing route in identity may not be overwritten"
if ident.RouteToCluster != "" && teleportClusterChanged {
return trace.AccessDenied("%s", overwriteDeniedMsg)
}
if ident.KubernetesCluster != "" && kubeClusterChanged {
return trace.AccessDenied("%s", overwriteDeniedMsg)
}
return nil
}
// singleCertHandler extracts routing information from base64-encoded URL
// parameters into the current auth user context and forwards the request back
// to the main router with the path prefix (and its embedded routing parameters)
// stripped.
func (f *Forwarder) singleCertHandler() httprouter.Handle {
return httplib.MakeHandlerWithErrorWriter(func(w http.ResponseWriter, req *http.Request, p httprouter.Params) (any, error) {
teleportCluster, kubeCluster, err := parseRouteFromPath(p)
if err != nil {
return nil, trace.Wrap(err)
}
userTypeI, err := authz.UserFromContext(req.Context())
if err != nil {
f.log.WarnContext(req.Context(), "error getting user from context", "error", err)
return nil, trace.AccessDenied("%s", accessDeniedMsg)
}
// Insert the extracted routing information from the path into the
// identity. Some implementation notes:
// - This still relies on RouteToCluster and KubernetesCluster identity
// fields, even though these fields are not part of the TLS identity
// when using path-based routing.
// - If the Teleport+Kube cluster names resolve to the local node, these
// values will be used directly in their proper handlers once the
// request is rewritten.
// - If the route resolves to a remote node, the identity is encoded (in
// JSON form) into forwarding headers using
// `auth.IdentityForwardingHeaders`. The destination node's auth
// middleware is configured to extract this identity (due to
// EnableCredentialsForwarding) and implicitly trusts this routing
// data, assuming the request originated from a proxy.
// - In either case, the destination node is ultimately responsible for
// authorizing the request, and routing information set in the
// identity should not be implicitly trusted. (This was ideally never
// the case, given access to resources could be revoked via roles
// before certs expired.)
var userType authz.IdentityGetter
switch o := userTypeI.(type) {
case authz.LocalUser:
if err := ensureRouteNotOverwritten(&o.Identity, teleportCluster, kubeCluster); err != nil {
return nil, trace.Wrap(err)
}
o.Identity.RouteToCluster = teleportCluster
o.Identity.KubernetesCluster = kubeCluster
userType = o
case authz.RemoteUser:
if err := ensureRouteNotOverwritten(&o.Identity, teleportCluster, kubeCluster); err != nil {
return nil, trace.Wrap(err)
}
o.Identity.RouteToCluster = teleportCluster
o.Identity.KubernetesCluster = kubeCluster
userType = o
default:
f.log.WarnContext(req.Context(), "Denying proxy access to unsupported user type", "user_type", logutils.TypeAttr(userTypeI))
return nil, trace.AccessDenied("%s", accessDeniedMsg)
}
ctx := authz.ContextWithUser(req.Context(), userType)
req = req.Clone(ctx)
path := p.ByName("path")
if !strings.HasPrefix(path, "/") {
path = "/" + path
}
req.URL.Path = path
req.URL.RawPath = ""
req.RequestURI = req.URL.RequestURI()
f.router.ServeHTTP(w, req)
return nil, nil
}, f.formatStatusResponseError)
}
/*
* 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 proxy
import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
"net"
"net/http"
"time"
"github.com/gravitational/trace"
semconv "go.opentelemetry.io/otel/semconv/v1.4.0"
oteltrace "go.opentelemetry.io/otel/trace"
"golang.org/x/net/http2"
"k8s.io/client-go/transport"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/healthcheck"
"github.com/gravitational/teleport/lib/kube/internal"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils"
)
// transportForRequest returns a transport that can be used to dial the next hop
// for the provided request using Impersonation.
func (f *Forwarder) transportForRequest(sess *clusterSession) (http.RoundTripper, error) {
transport, _, err := f.transportForRequestWithImpersonation(sess)
return transport, trace.Wrap(err)
}
// dialContextFunc is a context network dialer function that returns a network connection
type dialContextFunc func(context.Context, string, string) (net.Conn, error)
// transportForRequestWithImpersonation returns a transport that supports
// impersonation. This allows the client to reuse the same transport for all
// requests to the cluster in order to improve performance.
// The transport is cached in the forwarder so that it can be reused for future
// requests. If the transport is not cached, a new one is created and cached.
func (f *Forwarder) transportForRequestWithImpersonation(sess *clusterSession) (http.RoundTripper, *tls.Config, error) {
// If the session has a kube API credentials, it means that the next hop is
// a Kubernetes API server. In this case, we can use the provided credentials
// to dial the next hop directly and never cache the transport.
if sess.kubeAPICreds != nil {
// If agent is running in agent mode, get the transport from the configured cluster
// credentials.
return sess.kubeAPICreds.getTransport(), sess.kubeAPICreds.getTLSConfig(), nil
}
// If the cluster is remote, the key is the teleport cluster name.
// If the cluster is local, the key is the teleport cluster name and the kubernetes
// cluster name: <teleport-cluster-name>/<kubernetes-cluster-name>.
// If the cluster is local and scoped, the key is the teleport cluster name, the scope, and the
// kubernetes cluster name: <teleport-cluster-name>/<scope>/<kubernetes-cluster-name>.
key := transportCacheKey(sess)
t, err := utils.FnCacheGet(f.ctx, f.cachedTransport, key, func(ctx context.Context) (*cachedTransportEntry, error) {
var (
httpTransport http.RoundTripper
tlsConfig *tls.Config
err error
)
if sess.teleportCluster.isRemote {
// If the cluster is remote, create a new transport for the remote cluster.
httpTransport, tlsConfig, err = f.newRemoteClusterTransport(sess.teleportCluster.name)
} else if f.cfg.ReverseTunnelSrv != nil {
// If agent is running in proxy mode, create a new transport for the local cluster.
httpTransport, tlsConfig, err = f.newLocalClusterTransport(sess.kubeClusterName, sess.getClusterScope())
} else {
return nil, trace.BadParameter("no reverse tunnel server or credentials provided")
}
if err != nil {
return nil, trace.Wrap(err)
}
return &cachedTransportEntry{
transport: httpTransport,
tlsConfig: tlsConfig,
}, nil
})
if err != nil {
return nil, nil, trace.Wrap(err)
}
return t.transport, t.tlsConfig.Clone(), nil
}
// transportCacheKey returns a key used to cache transports.
// If the cluster is remote, the key is the teleport cluster name.
// If the cluster is local and scoped, the key is the teleport cluster name, the scope, and the
// kubernetes cluster name.
// If the cluster is local and unscoped, the key is the teleport cluster name and the kubernetes
// cluster name.
// The key is used to cache transports so that they can be reused for future requests.
// Each transport contains a custom dialer that is valid for a specific Teleport
// remote proxy or Teleport Kubernetes Services that serves the target cluster.
func transportCacheKey(sess *clusterSession) string {
if sess.teleportCluster.isRemote {
return fmt.Sprintf("%x", sess.teleportCluster.name)
}
if scope := sess.getClusterScope(); scope != "" {
return fmt.Sprintf("%x/%x/%x", sess.teleportCluster.name, scope, sess.kubeClusterName)
}
return fmt.Sprintf("%x/%x", sess.teleportCluster.name, sess.kubeClusterName)
}
// wrapTransport wraps the provided transport with the Kubernetes transport config
// if it is not nil.
func wrapTransport(rt http.RoundTripper, transportConfig *transport.Config) (http.RoundTripper, error) {
if transportConfig == nil {
return rt, nil
}
wrapped, err := transport.HTTPWrappersForConfig(transportConfig, rt)
if err != nil {
return nil, trace.Wrap(err)
}
return enforceCloseIdleConnections(wrapped, rt), nil
}
// newTransport creates a new [http.Transport] with the provided dialer and TLS
// config.
// The transport is configured to use a connection pool and to close idle
// connections after a timeout.
func newTransport(dial dialContextFunc, tlsConfig *tls.Config) *http.Transport {
return &http.Transport{
DialContext: dial,
TLSClientConfig: tlsConfig,
// Increase the size of the connection pool. This substantially improves the
// performance of Teleport under load as it reduces the number of TLS
// handshakes performed.
MaxIdleConns: defaults.HTTPMaxIdleConns,
MaxIdleConnsPerHost: defaults.HTTPMaxIdleConnsPerHost,
// IdleConnTimeout defines the maximum amount of time before idle connections
// are closed. Leaving this unset will lead to connections open forever and
// will cause memory leaks in a long running process.
IdleConnTimeout: defaults.HTTPIdleTimeout,
}
}
// newRemoteClusterTransport returns a new [http.Transport] (https://golang.org/pkg/net/http/#Transport)
// that can be used to dial Kubernetes Proxy in a remote Teleport cluster.
// The transport is configured to use a connection pool and to close idle
// connections after a timeout.
func (f *Forwarder) newRemoteClusterTransport(clusterName string) (http.RoundTripper, *tls.Config, error) {
// Tunnel is nil for a teleport process with "kubernetes_service" but
// not "proxy_service".
if f.cfg.ReverseTunnelSrv == nil {
return nil, nil, trace.BadParameter("this Teleport process can not dial Kubernetes endpoints in remote Teleport clusters; only proxy_service supports this, make sure a Teleport proxy is first in the request path")
}
// Dialer that will be used to dial the remote cluster via the reverse tunnel.
dialFn := f.remoteClusterDialer(clusterName)
tlsConfig, err := f.getTLSConfigForLeafCluster(clusterName)
if err != nil {
return nil, nil, trace.Wrap(err)
}
// Create a new HTTP/2 transport that will be used to dial the remote cluster.
h2Transport, err := newH2Transport(tlsConfig, dialFn)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return instrumentedRoundtripper(
f.cfg.KubeServiceType,
internal.NewImpersonatorRoundTripper(h2Transport),
), tlsConfig.Clone(), nil
}
// getTLSConfigForLeafCluster returns a TLS config with the Proxy certificate
// and the root CAs for the leaf cluster. Root proxy uses its own certificate
// to connect to the leaf proxy.
func (f *Forwarder) getTLSConfigForLeafCluster(clusterName string) (*tls.Config, error) {
ctx, cancel := context.WithTimeout(f.ctx, 5*time.Second)
defer cancel()
// Get the host CA for the target cluster from Auth to ensure we trust the
// leaf proxy certificate at the current time.
_, err := f.cfg.CachingAuthClient.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: clusterName,
}, false)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig := utils.TLSConfig(f.cfg.ConnTLSCipherSuites)
tlsConfig.GetClientCertificate = func(*tls.CertificateRequestInfo) (*tls.Certificate, error) {
tlsCert, err := f.cfg.GetConnTLSCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
return tlsCert, nil
}
tlsConfig.InsecureSkipVerify = true
tlsConfig.VerifyConnection = utils.VerifyConnectionWithRoots(func() (*x509.CertPool, error) {
pool, _, err := authclient.ClientCertPool(f.ctx, f.cfg.CachingAuthClient, clusterName, types.HostCA)
if err != nil {
return nil, trace.Wrap(err)
}
return pool, nil
})
return tlsConfig, nil
}
// remoteClusterDialer returns a dialer that can be used to dial Kubernetes Proxy
// in a remote Teleport cluster via the reverse tunnel.
func (f *Forwarder) remoteClusterDialer(clusterName string) dialContextFunc {
return func(ctx context.Context, _, _ string) (net.Conn, error) {
ctx, span := f.cfg.tracer.Start(
ctx,
"kube.Forwarder/remoteClusterDiater",
oteltrace.WithSpanKind(oteltrace.SpanKindClient),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("reverse_tunnel.Dial"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
targetCluster, err := f.cfg.ReverseTunnelSrv.Cluster(ctx, clusterName)
if err != nil {
return nil, trace.Wrap(err)
}
return targetCluster.DialTCP(reversetunnelclient.DialParams{
// Send a sentinel value to the remote cluster because this connection
// will be used to forward multiple requests to the remote cluster from
// different users.
// IP Pinning is based on the source IP address of the connection that
// we transport over HTTP headers so it's not affected.
From: &utils.NetAddr{AddrNetwork: "tcp", Addr: "0.0.0.0:0"},
// Proxy uses reverse tunnel dialer to connect to Kubernetes in a leaf cluster
// and the targetKubernetes cluster endpoint is determined from the identity
// encoded in the TLS certificate. We're setting the dial endpoint to a hardcoded
// `kube.teleport.cluster.local` value to indicate this is a Kubernetes proxy request
To: &utils.NetAddr{AddrNetwork: "tcp", Addr: reversetunnelclient.LocalKubernetes},
ConnType: types.KubeTunnel,
})
}
}
// newLocalClusterTransport returns a new [http.Transport] (https://golang.org/pkg/net/http/#Transport)
// that can be used to dial Kubernetes Service in a local Teleport cluster.
func (f *Forwarder) newLocalClusterTransport(kubeClusterName, scope string) (http.RoundTripper, *tls.Config, error) {
tlsConfig := utils.TLSConfig(f.cfg.ConnTLSCipherSuites)
tlsConfig.GetClientCertificate = func(*tls.CertificateRequestInfo) (*tls.Certificate, error) {
tlsCert, err := f.cfg.GetConnTLSCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
return tlsCert, nil
}
tlsConfig.InsecureSkipVerify = true
tlsConfig.VerifyConnection = utils.VerifyConnectionWithRoots(f.cfg.GetConnTLSRoots)
dialFn := f.localClusterDialer(kubeClusterName, scope)
// Create a new HTTP/2 transport that will be used to dial the remote cluster.
h2Transport, err := newH2Transport(tlsConfig, dialFn)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return instrumentedRoundtripper(
f.cfg.KubeServiceType,
internal.NewImpersonatorRoundTripper(h2Transport),
), tlsConfig.Clone(), nil
}
// localClusterDialer returns a dialer that can be used to dial Kubernetes Service
// in a local Teleport cluster using the reverse tunnel.
// The endpoints are fetched from the cached auth client and are shuffled
// to avoid hotspots.
func (f *Forwarder) localClusterDialer(kubeClusterName, scope string, opts ...contextDialerOption) dialContextFunc {
opt := contextDialerOptions{}
for _, o := range opts {
o(&opt)
}
return func(ctx context.Context, _, _ string) (net.Conn, error) {
ctx, span := f.cfg.tracer.Start(
ctx,
"kube.Forwarder/localClusterDiater",
oteltrace.WithSpanKind(oteltrace.SpanKindClient),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String(f.cfg.KubeServiceType),
semconv.RPCMethodKey.String("reverse_tunnel.Dial"),
semconv.RPCSystemKey.String("kube"),
),
)
defer span.End()
// Not a remote cluster and we have a reverse tunnel server.
// Use the local reversetunnel.Site which knows how to dial by serverID
// (for "kubernetes_service" connected over a tunnel) and falls back to
// direct dial if needed.
localCluster, err := f.cfg.ReverseTunnelSrv.Cluster(ctx, f.cfg.ClusterName)
if err != nil {
return nil, trace.Wrap(err)
}
kubeServers, err := f.getKubernetesServersForKubeCluster(ctx, kubeClusterName)
if err != nil {
return nil, trace.Wrap(err)
}
// Dial kube servers in the order of health status healthy, unknown, and unhealthy.
// Each health group is shuffled to distribute load.
// Unknown servers and unhealthy servers are still dialed
// in case health status changed since last check.
var errs []error
for server := range healthcheck.OrderByTargetHealthStatus(kubeServers) {
// Skip servers that don't match the requested scope.
if scope != "" && scopes.Compare(server.GetScope(), scope) != scopes.Equivalent {
continue
}
// Validate that the requested kube cluster is registered.
if server.GetCluster().GetName() != kubeClusterName || !opt.matches(server.GetHostID()) {
continue
}
// serverID is a unique identifier of the server in the cluster.
// It is a combination of the server's hostname and the cluster name.
// <host_id>.<cluster_name>
serverID := server.GetHostID() + "." + f.cfg.ClusterName
conn, err := localCluster.DialTCP(reversetunnelclient.DialParams{
// Send a sentinel value to the remote cluster because this connection
// will be used to forward multiple requests to the remote cluster from
// different users.
// IP Pinning is based on the source IP address of the connection that
// we transport over HTTP headers so it's not affected.
From: &utils.NetAddr{AddrNetwork: "tcp", Addr: "0.0.0.0:0"},
To: &utils.NetAddr{AddrNetwork: "tcp", Addr: server.GetHostname()},
ConnType: types.KubeTunnel,
ServerID: serverID,
ProxyIDs: server.GetProxyIDs(),
TargetScope: server.GetScope(),
})
if err == nil {
opt.collect(server.GetHostID())
return conn, nil
}
errs = append(errs, err)
}
if len(errs) > 0 {
return nil, trace.NewAggregate(errs...)
}
return nil, trace.NotFound("kubernetes cluster %q is not found in teleport cluster %q", kubeClusterName, f.cfg.ClusterName)
}
}
// newH2Transport creates a new HTTP/2 transport with ALPN support.
func newH2Transport(tlsConfig *tls.Config, dial dialContextFunc) (*http.Transport, error) {
tlsConfig = tlsConfig.Clone()
if tlsConfig == nil {
tlsConfig = &tls.Config{}
}
tlsConfig.NextProtos = []string{http2.NextProtoTLS, teleport.HTTPNextProtoTLS}
h2HTTPTransport := newTransport(dial, tlsConfig)
// Upgrade transport to h2 where HTTP_PROXY and HTTPS_PROXY
// envs are not take into account purposely.
if err := http2.ConfigureTransport(h2HTTPTransport); err != nil {
return nil, trace.Wrap(err)
}
return h2HTTPTransport, nil
}
// getTLSConfig returns TLS config required to connect to the next hop.
// If the current Kubernetes service serves the target cluster, it returns the
// Kubernetes API tls configuration.
// If the current service is a proxy and the next hop supports impersonation,
// it returns the proxy's TLS config.
// Otherwise, it requests a certificate from the auth server with the identity
// of the user that is requesting the connection embedded in the certificate.
// The boolean returned indicates whether the upstream server supports
// impersonation.
func (f *Forwarder) getTLSConfig(sess *clusterSession) (*tls.Config, bool, error) {
if sess.kubeAPICreds != nil {
return sess.kubeAPICreds.getTLSConfig(), false, nil
}
_, tlsConfig, err := f.transportForRequestWithImpersonation(sess)
return tlsConfig, err == nil, trace.Wrap(err)
}
// getContextDialerFunc returns a dialer function that can be used to connect
// to the next hop.
// If the next hop is a remote cluster, it returns a dialer that connects to
// the remote cluster proxy using the reverse tunnel server.
// If the next hop is a kubernetes service, it returns a dialer that connects
// to the first available kubernetes service.
// If the next hop is a local cluster, it returns a dialer that directly dials
// to the next hop.
func (f *Forwarder) getContextDialerFunc(s *clusterSession, opts ...contextDialerOption) dialContextFunc {
if s.kubeAPICreds != nil {
// If this is a kubernetes service, we need to connect to the kubernetes
// API server using a direct dialer.
return new(net.Dialer).DialContext
} else if s.teleportCluster.isRemote {
// If this is a remote cluster, we need to connect to the local proxy
// and then forward the connection to the remote cluster.
return f.remoteClusterDialer(s.teleportCluster.name)
} else if f.cfg.ReverseTunnelSrv != nil {
// If this is a local cluster, we need to connect to the remote proxy
// and then forward the connection to the local cluster.
return f.localClusterDialer(s.kubeClusterName, s.getClusterScope(), opts...)
}
return new(net.Dialer).DialContext
}
// contextDialerOptions is a set of options that can be used to filter
// the hosts that the dialer connects to.
type contextDialerOptions struct {
hostIDFilter string
collectHostID *string
}
// matches returns true if the host matches the hostID of the dialer options or
// if the dialer hostID is empty.
func (c *contextDialerOptions) matches(hostID string) bool {
return c.hostIDFilter == "" || c.hostIDFilter == hostID
}
// collect sets the hostID that the dialer connected to if collectHostID is not nil.
func (c *contextDialerOptions) collect(hostID string) {
if c.collectHostID != nil {
*c.collectHostID = hostID
}
}
// contextDialerOption is a functional option for the contextDialerOptions.
type contextDialerOption func(*contextDialerOptions)
// withTargetHostID is a functional option that sets the hostID of the dialer.
// If the hostID is empty, the dialer will connect to the first available host.
// If the hostID is not empty, the dialer will connect to the host with the
// specified hostID. If that host is not available, the dialer will return an
// error.
func withTargetHostID(hostID string) contextDialerOption {
return func(o *contextDialerOptions) {
o.hostIDFilter = hostID
}
}
// withHostIDCollection is a functional option that sets the hostID of the dialer
// to the provided pointer.
func withHostIDCollection(hostID *string) contextDialerOption {
return func(o *contextDialerOptions) {
o.collectHostID = hostID
}
}
/*
* 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 proxy
import (
"bytes"
"io"
"net/http"
"path"
"slices"
"strings"
"github.com/gravitational/trace"
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
"k8s.io/apimachinery/pkg/runtime/schema"
"k8s.io/apimachinery/pkg/runtime/serializer"
utilnet "k8s.io/apimachinery/pkg/util/net"
"github.com/gravitational/teleport/api/types"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/kube/proxy/responsewriters"
)
// metaResource wraps the various representations of a Kubernetes resource.
type metaResource struct {
resourceDefinition *metav1.APIResource // Resource definition data from the schema.
requestedResource apiResource // User input, based on URL.
verb string // Verb of the user request.
isList bool
// unsupportedResource is set when the requested kind is unknown to a healthy
// discovery cache and so can't be matched against kubernetes_resources rules.
// The forwarder denies such requests instead of forwarding them unenforced.
unsupportedResource bool
}
func (mr *metaResource) isClusterWideResource() bool {
if mr == nil {
return false
}
return mr.resourceDefinition != nil && !mr.resourceDefinition.Namespaced
}
func (mr *metaResource) rbacResource() *types.KubernetesResource {
if mr == nil || mr.resourceDefinition == nil {
return nil
}
return &types.KubernetesResource{
Kind: mr.resourceDefinition.Name,
Namespace: mr.requestedResource.namespace,
Name: mr.requestedResource.resourceName,
Verbs: []string{mr.verb},
APIGroup: mr.requestedResource.apiGroup,
}
}
// ephemeralContainersResourceKind is the pods subresource used to add an
// ephemeral container to a running pod (the endpoint `kubectl debug` targets).
const ephemeralContainersResourceKind = "pods/ephemeralcontainers"
// requiredRBACResources returns the set of Kubernetes resources the caller must be granted to perform this request.
// Most requests need a single (resource, verb) tuple, the one returned by rbacResource.
// Adding an ephemeral container runs code in the target pod the same way exec does,
// and the request also mutates the pod object (it is a patch/update),
// so it requires both the exec verb and the mutation verb the HTTP method maps to.
// Returning both keeps either verb on its own from being enough to add an ephemeral container.
func (mr *metaResource) requiredRBACResources() []types.KubernetesResource {
base := mr.rbacResource()
if base == nil {
return nil
}
if mr.requestedResource.resourceKind != ephemeralContainersResourceKind {
return []types.KubernetesResource{*base}
}
execResource := *base
execResource.Verbs = []string{types.KubeVerbExec}
return []types.KubernetesResource{*base, execResource}
}
// apiResource represents the resource requested by the user.
type apiResource struct {
apiGroup string
apiGroupVersion string
namespace string
resourceKind string
resourceName string
skipEvent bool
isWatch bool
isProxyVerb bool
}
// parseResourcePath does best-effort parsing of a Kubernetes API request path.
// All fields of the returned apiResource may be empty.
//
// TODO(jakealti): reuse k8s.io/apiserver request.RequestInfoFactory here instead of re-implementing it.
func parseResourcePath(p string) (apiResource, error) {
// Kubernetes API reference: https://kubernetes.io/docs/reference/kubernetes-api/
// Let's try to parse this. Here be dragons!
//
// URLs have a prefix that defines an "API group":
// - /api/v1/ - the special "" API group (e.g. pods, secrets, etc. belong here)
// - /apis/{group}/{version} - the other properly named groups (e.g. apps/v1 or rbac.authorization.k8s.io/v1beta1)
//
// After the prefix, we have the resource info:
// - /namespaces/{namespace}/{resource kind}/{resource name} for namespaced resources
// - turns out, namespace is optional when you query across all
// namespaces (e.g. /api/v1/pods to get pods in all namespaces)
// - /{resource kind}/{resource name} for cluster-scoped resources (e.g. namespaces or nodes)
//
// If {resource name} is missing, the request refers to all resources of
// that kind (e.g. list all pods).
//
// There can be more items after {resource name} (a "subresource"), like
// pods/foo/exec, but the depth is arbitrary (e.g.
// /api/v1/namespaces/{namespace}/pods/{name}/proxy/{path})
//
// And the cherry on top - watch endpoints, e.g.
// /api/v1/watch/namespaces/{namespace}/pods/{name}
// for live updates on resources (specific resources or all of one kind)
var r apiResource
// Cleaning below resolves "." and ".." segments, which changes which resource the path names.
// The raw path is what gets forwarded, so parsing one and forwarding the other
// would authorize a name the cluster never resolves.
for segment := range strings.SplitSeq(p, "/") {
if segment == "." || segment == ".." {
return r, trace.BadParameter("kubernetes request path must not contain %q segments", segment)
}
}
// Clean up the path and make it absolute.
p = path.Clean(p)
if !path.IsAbs(p) {
p = "/" + p
}
parts := strings.Split(p, "/")
switch {
// Core API group has a "special" URL prefix /api/v1/.
case len(parts) >= 3 && parts[1] == "api" && parts[2] == "v1":
r.apiGroup = ""
r.apiGroupVersion = parts[2]
parts = parts[3:]
// Other API groups have URL prefix /apis/{group}/{version}.
case len(parts) >= 4 && parts[1] == "apis":
r.apiGroup, r.apiGroupVersion = parts[2], parts[3]
parts = parts[4:]
case len(parts) >= 2 && (parts[1] == "api" || parts[1] == "apis"):
// /api or /apis.
// This is part of API discovery. Don't emit to audit log to reduce
// noise.
r.skipEvent = true
return r, nil
default:
// Doesn't look like a k8s API path, return empty result.
return r, nil
}
// Special verb endpoints carry the verb in a segment ahead of the resource path.
// The API server consumes that segment and parses the rest as a normal resource path, so we do the same.
// See specialVerbs in https://github.com/kubernetes/apiserver/blob/master/pkg/endpoints/request/requestinfo.go
if len(parts) > 1 {
switch parts[0] {
case "watch":
r.isWatch = true
parts = parts[1:]
case "proxy":
r.isProxyVerb = true
parts = parts[1:]
}
}
switch len(parts) {
case 0:
// e.g. /apis/apps/v1
// This is part of API discovery. Don't emit to audit log to reduce
// noise.
r.skipEvent = true
return r, nil
case 1:
// e.g. /api/v1/pods - list pods in all namespaces
r.resourceKind = parts[0]
case 2:
// e.g. /api/v1/clusterroles/{name} - read a cluster-level resource
r.resourceKind = parts[0]
r.resourceName = parts[1]
case 3:
if parts[0] == "namespaces" {
// e.g. /api/v1/namespaces/{namespace}/pods - list pods in a
// specific namespace
r.namespace = parts[1]
r.resourceKind = parts[2]
} else {
// e.g. /apis/apiregistration.k8s.io/v1/apiservices/{name}/status
kind := append([]string{parts[0]}, parts[2:]...)
r.resourceKind = strings.Join(kind, "/")
r.resourceName = parts[1]
}
default:
// e.g. /api/v1/namespaces/{namespace}/pods/{name} - get a specific pod
// or /api/v1/namespaces/{namespace}/pods/{name}/exec - exec command in a pod
if parts[0] == "namespaces" {
r.namespace = parts[1]
kind := append([]string{parts[2]}, parts[4:]...)
r.resourceKind = strings.Join(kind, "/")
r.resourceName = parts[3]
if len(parts) > 4 && parts[4] == "proxy" {
r.resourceName = stripProxyNamePortScheme(r.resourceName)
}
} else {
// e.g. /api/v1/nodes/{name}/proxy/{path}
kind := append([]string{parts[0]}, parts[2:]...)
r.resourceKind = strings.Join(kind, "/")
r.resourceName = parts[1]
if len(parts) > 2 && parts[2] == "proxy" {
r.resourceName = stripProxyNamePortScheme(r.resourceName)
}
}
}
// The proxy special verb takes no subresource.
// Kubernetes apiserver stops parsing at the name and hands every segment after it to the proxied backend.
// Drop them so the kind we record and report is the one the API server resolved.
if r.isProxyVerb {
r.resourceKind = getResourceFromAPIResource(r.resourceKind)
}
// The core API accepts [scheme:]name[:port] in the name segment of its pods/services/nodes proxy endpoints,
// so the special verb form has to be normalized the same way the subresource form above is.
// Otherwise a rule naming the resource stops matching once a scheme or port is supplied.
if r.isProxyVerb && r.apiGroup == "" {
switch r.resourceKind {
case "pods", "services", "nodes":
r.resourceName = stripProxyNamePortScheme(r.resourceName)
}
}
return r, nil
}
// stripProxyNamePortScheme extracts the bare resource name from the [scheme:]name[:port] segment that
// the Kubernetes API server accepts on pods/{name}/proxy, services/{name}/proxy, and nodes/{name}/proxy.
func stripProxyNamePortScheme(segment string) string {
_, name, _, valid := utilnet.SplitSchemeNamePort(segment)
if !valid {
return segment
}
return name
}
func (r apiResource) populateEvent(e *apievents.KubeRequest) {
e.ResourceAPIGroup = path.Join(r.apiGroup, r.apiGroupVersion)
e.ResourceNamespace = r.namespace
e.ResourceKind = r.resourceKind
e.ResourceName = r.resourceName
}
// allowedResourcesKey is a key used to identify a resource in the allowedResources map.
type allowedResourcesKey struct {
apiGroup string
resourceKind string
}
type rbacSupportedResources map[allowedResourcesKey]metav1.APIResource
// getResourceWithKey returns the teleport resource kind for a given resource key if
// it exists, otherwise returns an empty string.
func (r rbacSupportedResources) getResource(apiGroup, resourceKind string) (metav1.APIResource, bool) {
k := allowedResourcesKey{
apiGroup: apiGroup,
resourceKind: getResourceFromAPIResource(resourceKind),
}
out, ok := r[k]
return out, ok
}
func (r rbacSupportedResources) getTeleportResourceKindFromAPIResource(api apiResource) (string, bool) {
resource := getResourceFromAPIResource(api.resourceKind)
resourceType, ok := r[allowedResourcesKey{apiGroup: api.apiGroup, resourceKind: resource}]
return resourceType.Kind, ok
}
// getResourceFromRequest returns a KubernetesResource if the user tried to access
// a specific endpoint that Teleport support resource filtering. Otherwise, returns nil.
func getResourceFromRequest(req *http.Request, kubeDetails *kubeDetails) (metaResource, error) {
apiResource, err := parseResourcePath(req.URL.Path)
if err != nil {
return metaResource{}, trace.Wrap(err)
}
out := metaResource{
requestedResource: apiResource,
verb: apiResource.getVerb(req),
}
if kubeDetails == nil {
return out, nil
}
// Surface an offline-cluster error before doing any work.
if _, _, err := kubeDetails.getClusterSupportedResources(); err != nil {
return out, trace.Wrap(err)
}
// Discovery / health endpoints (e.g. /api, /apis/<group>/<version>) carry no
// concrete resource kind and aren't subject to kubernetes_resources rules;
// let them through so clients can complete API discovery.
if apiResource.resourceKind == "" {
return out, nil
}
resource, found := kubeDetails.resolveResource(apiResource.apiGroup, apiResource.apiGroupVersion, apiResource.resourceKind)
if !found {
// Still unknown after a targeted discovery: the kind isn't served by the
// cluster. Flag it so the forwarder denies the request rather than
// forwarding it with kubernetes_resources rules unenforced.
out.unsupportedResource = true
return out, nil
}
out.resourceDefinition = &resource
if apiResource.resourceName == "" && !slices.Contains([]string{types.KubeVerbCreate, types.KubeVerbProxy}, out.verb) {
// A missing name means the request targets the whole collection: a list.
// Create and proxy are the exceptions. Create has the name in the body, proxy has none.
out.isList = true
return out, nil
}
if apiResource.resourceName == "" && out.verb == types.KubeVerbCreate {
// If the request is a create request, extract the resource name from the request body.
// Re-read the codecs: resolveResource above may have just discovered this CRD and
// rebuilt them, so decode against a scheme that knows the new kind.
codecFactory, _, err := kubeDetails.getClusterSupportedResources()
if err != nil {
return out, trace.Wrap(err)
}
resourceName, err := extractResourceNameFromPostRequest(req, codecFactory, kubeDetails.getObjectGVK(apiResource))
if err != nil {
return out, trace.Wrap(err)
}
apiResource.resourceName = resourceName
out.requestedResource = apiResource
}
return out, nil
}
// extractResourceNameFromPostRequest extracts the resource name from a POST body.
// It reads the full body - required because data can be proto encoded -
// and decodes it into a Kubernetes object. It then extracts the resource name
// from the object.
func extractResourceNameFromPostRequest(
req *http.Request,
codecs *serializer.CodecFactory,
defaults *schema.GroupVersionKind,
) (string, error) {
if req.Body == nil {
return "", trace.BadParameter("request body is empty")
}
negotiator := newClientNegotiator(codecs)
_, decoder, err := newEncoderAndDecoderForContentType(
responsewriters.GetContentTypeHeader(req.Header),
negotiator,
)
if err != nil {
return "", trace.Wrap(err)
}
newBody := bytes.NewBuffer(make([]byte, 0, 2048))
if _, err := io.Copy(newBody, req.Body); err != nil {
return "", trace.Wrap(err)
}
if err := req.Body.Close(); err != nil {
return "", trace.Wrap(err)
}
// The body is replaced with a replayable reader, and [http.Request.GetBody] is
// set so the upstream transport can retry the request after a GOAWAY without
// failing on the unrewindable network-side body.
// See https://github.com/gravitational/teleport/issues/65611
bodyBytes := newBody.Bytes()
req.Body = io.NopCloser(bytes.NewReader(bodyBytes))
req.GetBody = func() (io.ReadCloser, error) {
return io.NopCloser(bytes.NewReader(bodyBytes)), nil
}
req.ContentLength = int64(len(bodyBytes))
// decode memory rw body.
obj, err := decodeAndSetGVK(decoder, bodyBytes, defaults)
if err != nil {
return "", trace.Wrap(err)
}
namer, ok := obj.(kubeObjectInterface)
if !ok {
return "", trace.BadParameter("object %T does not implement kubeObjectInterface", obj)
}
return namer.GetName(), nil
}
// getResourceFromAPIResource returns the resource kind from the api resource.
// If the resource kind contains sub resources (e.g. pods/exec), it returns the
// resource kind without the subresource.
func getResourceFromAPIResource(resourceKind string) string {
if idx := strings.Index(resourceKind, "/"); idx != -1 {
return resourceKind[:idx]
}
return resourceKind
}
// splitResourceSubresource splits a resourceKind into its base resource and
// the first subresource segment. The trailing path (if any) is discarded.
// Examples:
//
// "pods" -> "pods", ""
// "pods/exec" -> "pods", "exec"
// "pods/proxy/8080" -> "pods", "proxy"
// "nodes/proxy/foo/bar" -> "nodes", "proxy"
func splitResourceSubresource(resourceKind string) (base, sub string) {
parts := strings.SplitN(resourceKind, "/", 3)
if len(parts) < 2 {
return resourceKind, ""
}
return parts[0], parts[1]
}
// isKubeWatchRequest returns true if the request is a watch request.
func isKubeWatchRequest(req *http.Request, r apiResource) bool {
if values := req.URL.Query()["watch"]; len(values) > 0 {
switch strings.ToLower(values[0]) {
case "false", "0":
default:
return true
}
}
return r.isWatch
}
func (r apiResource) getVerb(req *http.Request) string {
if r.isProxyVerb {
return types.KubeVerbProxy
}
verb := ""
isWatch := isKubeWatchRequest(req, r)
switch r.resourceKind {
case "pods/exec", "pods/attach":
verb = types.KubeVerbExec
case "pods/portforward":
verb = types.KubeVerbPortForward
default:
if base, sub := splitResourceSubresource(r.resourceKind); sub == "proxy" {
switch base {
case "pods", "services", "nodes":
return types.KubeVerbProxy
}
}
switch req.Method {
case http.MethodPost:
verb = types.KubeVerbCreate
case http.MethodGet, http.MethodHead, http.MethodOptions:
switch {
case isWatch:
return types.KubeVerbWatch
case r.resourceName == "":
return types.KubeVerbList
default:
return types.KubeVerbGet
}
case http.MethodPut:
verb = types.KubeVerbUpdate
case http.MethodPatch:
verb = types.KubeVerbPatch
case http.MethodDelete:
switch {
case r.resourceName != "":
verb = types.KubeVerbDelete
default:
verb = types.KubeVerbDeleteCollection
}
default:
verb = ""
}
}
return verb
}
/*
* 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 proxy
import (
"context"
"maps"
"slices"
"sync"
"time"
"github.com/gravitational/trace"
presencev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/presence/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/utils"
)
// startReconciler starts reconciler that registers/unregisters proxied
// kubernetes clusters according to the up-to-date list of kube_cluster resources.
func (s *TLSServer) startReconciler(ctx context.Context) (err error) {
if len(s.ResourceMatchers) == 0 || s.KubeServiceType != KubeService {
s.log.DebugContext(ctx, "Not initializing Kube Cluster resource watcher")
return nil
}
s.reconciler, err = services.NewReconciler(services.ReconcilerConfig[types.KubeCluster]{
Matcher: s.matcher,
GetCurrentResources: s.getResources,
CompareResources: func(kc1, kc2 types.KubeCluster) int {
return services.EqualFromBool(kc1.IsEqual(kc2))
},
GetNewResources: s.monitoredKubeClusters.get,
OnCreate: s.onCreate,
OnUpdate: s.onUpdate,
OnDelete: s.onDelete,
Logger: s.log.With("kind", types.KindKubernetesCluster),
})
if err != nil {
return trace.Wrap(err)
}
go func() {
// reconcileTicker is used to force reconciliation when the watcher was
// previously informed that a `kube_cluster` resource exists/changed but the
// creation/update operation failed - e.g. login to AKS/EKS clusters can
// fail due to missing permissions.
// Once this happens, the state of the resource watcher won't change until
// a new update operation is triggered (which can take a lot of time).
// This results in the service not being able to enroll the failing cluster,
// even if the original issue was already fixed because we won't run reconciliation again.
// We force the reconciliation to make sure we don't drift from watcher state if
// the issue was fixed.
reconcileTicker := time.NewTicker(2 * time.Minute)
defer reconcileTicker.Stop()
for {
select {
case <-reconcileTicker.C:
if err := s.reconciler.Reconcile(ctx); err != nil {
s.log.ErrorContext(ctx, "Failed to reconcile", "error", err)
}
case <-s.reconcileCh:
if err := s.reconciler.Reconcile(ctx); err != nil {
s.log.ErrorContext(ctx, "Failed to reconcile", "error", err)
} else if s.OnReconcile != nil {
s.OnReconcile(s.fwd.kubeClusters())
}
case <-ctx.Done():
s.log.DebugContext(ctx, "Reconciler done")
return
}
}
}()
return nil
}
// startKubeClusterResourceWatcher starts watching changes to Kube Clusters resources and
// registers/unregisters the proxied Kube Cluster accordingly.
func (s *TLSServer) startKubeClusterResourceWatcher(ctx context.Context) (*services.GenericWatcher[types.KubeCluster, readonly.KubeCluster], error) {
if len(s.ResourceMatchers) == 0 || s.KubeServiceType != KubeService {
s.log.DebugContext(ctx, "Not initializing Kube Cluster resource watcher")
return nil, nil
}
s.log.DebugContext(ctx, "Initializing Kube Cluster resource watcher")
watcher, err := services.NewKubeClusterWatcher(ctx, services.KubeClusterWatcherConfig{
ResourceWatcherConfig: services.ResourceWatcherConfig{
Component: s.Component,
Logger: s.log,
Client: s.AccessPoint,
},
KubernetesClusterGetter: s.AccessPoint,
// dynamically registered clusters may contain secrets which are necessary in order
// for the agent to connect to the cluster.
LoadSecrets: true,
})
if err != nil {
return nil, trace.Wrap(err)
}
go func() {
defer watcher.Close()
for {
select {
case clusters := <-watcher.ResourcesC:
s.monitoredKubeClusters.setResources(clusters)
select {
case s.reconcileCh <- struct{}{}:
case <-ctx.Done():
return
}
case <-ctx.Done():
s.log.DebugContext(ctx, "Kube Cluster resource watcher done")
return
}
}
}()
return watcher, nil
}
func (s *TLSServer) getResources() map[string]types.KubeCluster {
return utils.FromSlice(s.fwd.kubeClusters(), types.KubeCluster.GetName)
}
func (s *TLSServer) onCreate(ctx context.Context, cluster types.KubeCluster) error {
return s.registerKubeCluster(ctx, cluster)
}
func (s *TLSServer) onUpdate(ctx context.Context, cluster, _ types.KubeCluster) error {
return s.updateKubeCluster(ctx, cluster)
}
func (s *TLSServer) onDelete(ctx context.Context, cluster types.KubeCluster) error {
return s.unregisterKubeCluster(ctx, cluster, false)
}
func (s *TLSServer) matcher(cluster types.KubeCluster) bool {
return services.MatchResourceLabels(s.ResourceMatchers, cluster.GetAllLabels())
}
// monitoredKubeClusters is a collection of clusters from different sources
// like configuration file and dynamic resources.
//
// It's updated by respective watchers and is used for reconciling with the
// currently proxied clusters.
type monitoredKubeClusters struct {
// static are clusters from the agent's YAML configuration.
static types.KubeClusters
// resources are clusters created via CLI or API.
resources types.KubeClusters
// mu protects access to the fields.
mu sync.Mutex
}
func (m *monitoredKubeClusters) setResources(clusters types.KubeClusters) {
m.mu.Lock()
defer m.mu.Unlock()
m.resources = clusters
}
func (m *monitoredKubeClusters) get() map[string]types.KubeCluster {
m.mu.Lock()
defer m.mu.Unlock()
return utils.FromSlice(append(m.static, m.resources...), types.KubeCluster.GetName)
}
func (s *TLSServer) buildClusterDetailsConfigForCluster(cluster types.KubeCluster) clusterDetailsConfig {
return clusterDetailsConfig{
azureClients: s.azureClients,
gcpClients: s.gcpClients,
awsCloudClients: s.awsClients,
cluster: cluster,
log: s.log,
checker: s.CheckImpersonationPermissions,
resourceMatchers: s.ResourceMatchers,
clock: s.Clock,
component: s.KubeServiceType,
}
}
func (s *TLSServer) registerKubeCluster(ctx context.Context, cluster types.KubeCluster) error {
clusterDetails, err := newClusterDetails(
ctx,
s.buildClusterDetailsConfigForCluster(cluster),
)
if err != nil {
return trace.Wrap(err)
}
s.fwd.upsertKubeDetails(cluster.GetName(), clusterDetails)
return trace.Wrap(s.startHeartbeatAndHealthCheck(cluster))
}
func (s *TLSServer) updateKubeCluster(ctx context.Context, cluster types.KubeCluster) error {
clusterDetails, err := newClusterDetails(
ctx,
s.buildClusterDetailsConfigForCluster(cluster),
)
if err != nil {
return trace.Wrap(err)
}
s.fwd.upsertKubeDetails(cluster.GetName(), clusterDetails)
return nil
}
// unregisterKubeCluster unregisters the proxied Kube Cluster from the agent.
// This function is called when the dynamic cluster is deleted/no longer match
// the agent's resource matcher or when the agent is shutting down.
func (s *TLSServer) unregisterKubeCluster(ctx context.Context, cluster types.KubeCluster, isShutdown bool) error {
var errs []error
if err := s.stopHeartbeatAndHealthCheck(cluster); err != nil {
errs = append(errs, err)
}
clusterName := cluster.GetName()
s.fwd.removeKubeDetails(clusterName)
// A child process can be forked to upgrade the Teleport binary. The child
// will take over the heartbeats so do NOT delete them in that case.
shouldDeleteOnShutdown := services.ShouldDeleteServerHeartbeatsOnShutdown(ctx)
sender, ok := s.TLSServerConfig.InventoryHandle.GetSender()
if ok {
// Manual deletion per cluster is only required if the auth server
// doesn't support actively cleaning up database resources when the
// inventory control stream is terminated during shutdown.
if capabilities := sender.Hello().GetCapabilities(); capabilities != nil {
shouldDeleteOnShutdown = shouldDeleteOnShutdown && !capabilities.GetKubernetesCleanup()
}
}
if !isShutdown || shouldDeleteOnShutdown {
if err := s.deleteKubernetesServer(ctx, clusterName); err != nil {
errs = append(errs, err)
}
}
// close active sessions before returning.
s.fwd.mu.Lock()
// collect all sessions to avoid holding the lock while closing them
sessions := slices.Collect(maps.Values(s.fwd.sessions))
s.fwd.mu.Unlock()
// close active sessions
for _, sess := range sessions {
if sess.ctx.kubeClusterName == clusterName {
// TODO(tigrato): check if we should send errors to each client
if err := sess.Close(); err != nil {
errs = append(errs, err)
}
}
}
return trace.NewAggregate(errs...)
}
// deleteKubernetesServer deletes kubernetes server for the specified cluster.
func (s *TLSServer) deleteKubernetesServer(ctx context.Context, name string) error {
err := s.AuthClient.DeleteKubeServer(ctx, presencev1.DeleteKubeServerRequest_builder{
Scope: s.GetScope(),
HostId: s.HostID,
Name: name,
}.Build())
if err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err)
}
return nil
}
/*
* 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 multiplexer implements SSH and TLS multiplexing
// on the same listener
//
// mux, _ := multiplexer.New(Config{Listener: listener})
// mux.SSH() // returns listener getting SSH connections
// mux.TLS() // returns listener getting TLS connections
package multiplexer
import (
"bufio"
"bytes"
"context"
"crypto"
"crypto/sha256"
"errors"
"io"
"log/slog"
"net"
"slices"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/jwt"
"github.com/gravitational/teleport/lib/loglimit"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
var (
// ErrBadIP is returned when there's a problem with client source or destination IP address
ErrBadIP = &trace.BadParameterError{Message: "client source and destination addresses should be valid same TCP version non-nil IP addresses"}
// ErrDowngradeDst is returned when attempting to downgrade an IPv6 destination instead of an IPv6 source
ErrDowngradeDst = &trace.BadParameterError{Message: "only client source addresses can be downgraded to IPv4, downgrading destination addresses is not supported"}
)
// Start of class E IPv4 CIDR range
const classEPrefix byte = 240
// PROXYProtocolMode controls behavior related to unsigned PROXY protocol headers.
// Possible values:
// - 'on': one PROXY header is accepted and required per incoming connection.
// - 'off': no PROXY headers are allows, otherwise connection is rejected.
// If unspecified - one PROXY header is allowed, but not required. Connection is marked with source port set to 0
// and IP pinning will not be allowed. It is supposed to be used only as default mode for test setups.
// In production you should always explicitly set the mode based on your network setup - if you have L4 load balancer
// with enabled PROXY protocol in front of Teleport you should set it to 'on', if you don't have it, set it to 'off'
type PROXYProtocolMode string
const (
PROXYProtocolOn PROXYProtocolMode = "on"
PROXYProtocolOff PROXYProtocolMode = "off"
PROXYProtocolUnspecified PROXYProtocolMode = ""
)
// CertAuthorityGetter allows to get cluster's host CA for verification of signed PROXY headers.
// We define our own version to not create dependency on the 'services' package, which causes circular references
type CertAuthorityGetter = func(ctx context.Context, id types.CertAuthID, loadKeys bool) (types.CertAuthority, error)
type (
// PreDetectFunc is used in [Mux]'s [Config] as the PreDetect hook.
PreDetectFunc = func(net.Conn) (PostDetectFunc, error)
// PostDetectFunc is optionally returned by a [PreDetectFunc].
PostDetectFunc = func(*Conn) net.Conn
)
// Config is a multiplexer config
type Config struct {
// Listener is listener to multiplex connection on
Listener net.Listener
// Context is a context to signal stops, cancellations
Context context.Context
// DetectTimeout is a timeout applied to the whole detection phase of the
// connection, set to defaults.ReadHeadersTimeout if unspecified
DetectTimeout time.Duration
// Clock is a clock to override in tests, set to real time clock
// by default
Clock clockwork.Clock
// PROXYProtocolMode controls behavior related to unsigned PROXY protocol headers.
PROXYProtocolMode PROXYProtocolMode
// PROXYAllowDowngrade controls IPv6 downgrade to pseudo IPv4 in PROXY headers
PROXYAllowDowngrade bool
// SuppressUnexpectedPROXYWarning makes multiplexer not issue warnings if it receives PROXY
// line when running in PROXYProtocolMode=PROXYProtocolUnspecified
SuppressUnexpectedPROXYWarning bool
// ID is an identifier used for debugging purposes
ID string
// CertAuthorityGetter is used to get CA to verify singed PROXY headers sent internally by teleport
CertAuthorityGetter CertAuthorityGetter
// LocalClusterName set the local cluster for the multiplexer, it's used in PROXY headers verification.
LocalClusterName string
// IgnoreSelfConnections is used for tests, it makes multiplexer ignore the fact that it's self
// connection (coming from same IP as the listening address) when deciding if it should drop connection with
// missing required PROXY header. This is needed since all connections in tests are self connections.
IgnoreSelfConnections bool
// PreDetect, if set, is called on each incoming connection before protocol
// detection; the returned [PostDetectFunc] (if any) will then be called
// after protocol detection, and will have the ability to modify or wrap the
// [*Conn] before it's passed to the listener; if the PostDetectFunc returns
// a nil [net.Conn], the connection will not be handled any further by the
// multiplexer, and it's the responsibility of the PostDetectFunc to arrange
// for it to be eventually closed.
PreDetect PreDetectFunc
}
// CheckAndSetDefaults verifies configuration and sets defaults
func (c *Config) CheckAndSetDefaults() error {
if c.Listener == nil {
return trace.BadParameter("missing parameter Listener")
}
if c.Context == nil {
c.Context = context.TODO()
}
if c.DetectTimeout == 0 {
c.DetectTimeout = defaults.ReadHeadersTimeout
}
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
return nil
}
// New returns a new instance of multiplexer
func New(cfg Config) (*Mux, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
ctx, cancel := context.WithCancel(cfg.Context)
logLimiter, err := loglimit.New(loglimit.Config{
MessageSubstrings: errorSubstrings,
Handler: slog.Default().Handler(),
})
if err != nil {
cancel()
return nil, trace.Wrap(err)
}
waitContext, waitCancel := context.WithCancel(context.TODO())
return &Mux{
logger: slog.With(teleport.ComponentKey, teleport.Component("mx", cfg.ID)),
Config: cfg,
context: ctx,
cancel: cancel,
waitContext: waitContext,
waitCancel: waitCancel,
sampledLogger: slog.New(logLimiter).With(teleport.ComponentKey, teleport.Component("mx", cfg.ID)),
}, nil
}
// Mux supports having both SSH and TLS on the same listener socket
type Mux struct {
sync.RWMutex
logger *slog.Logger
Config
sshListener *Listener
tlsListener *Listener
dbListener *Listener
httpListener *Listener
context context.Context
cancel context.CancelFunc
waitContext context.Context
waitCancel context.CancelFunc
// sampledLogger is a logger responsible for deduplicating multiplexer errors
// (over a 1min window) that occur when detecting the types of new connections.
// This ensures that health checkers / malicious actors cannot overpower /
// pollute the logs with warnings when such connections are invalid or unknown
// to the multiplexer.
sampledLogger *slog.Logger
}
// SSH returns listener that receives SSH connections
func (m *Mux) SSH() net.Listener {
m.Lock()
defer m.Unlock()
if m.sshListener == nil {
m.sshListener = newListener(m.context, m.Config.Listener.Addr())
}
return m.sshListener
}
// TLS returns listener that receives TLS connections
func (m *Mux) TLS() net.Listener {
m.Lock()
defer m.Unlock()
if m.tlsListener == nil {
m.tlsListener = newListener(m.context, m.Config.Listener.Addr())
}
return m.tlsListener
}
// DB returns listener that receives database connections
func (m *Mux) DB() net.Listener {
m.Lock()
defer m.Unlock()
if m.dbListener == nil {
m.dbListener = newListener(m.context, m.Config.Listener.Addr())
}
return m.dbListener
}
// HTTP returns listener that receives plain HTTP connections
func (m *Mux) HTTP() net.Listener {
m.Lock()
defer m.Unlock()
if m.httpListener == nil {
m.httpListener = newListener(m.context, m.Config.Listener.Addr())
}
return m.httpListener
}
func (m *Mux) closeListener() {
m.Lock()
defer m.Unlock()
// propagate close signal to other listeners
m.cancel()
if m.Listener == nil {
return
}
m.Listener.Close()
}
// Close closes listener
func (m *Mux) Close() error {
m.closeListener()
return nil
}
// Wait waits until listener shuts down and stops accepting new connections
// this is to workaround issue https://github.com/golang/go/issues/10527
// in tests
func (m *Mux) Wait() {
<-m.waitContext.Done()
}
// Serve is a blocking function that serves on the listening socket
// and accepts requests. Every request is served in a separate goroutine
func (m *Mux) Serve() error {
m.logger.DebugContext(m.context, "Starting serving MUX", "listen_addr", m.Config.Listener.Addr())
defer m.waitCancel()
for {
conn, err := m.Listener.Accept()
if err == nil {
if tcpConn, ok := conn.(*net.TCPConn); ok {
tcpConn.SetKeepAlive(true)
tcpConn.SetKeepAlivePeriod(3 * time.Minute)
}
go m.detectAndForward(conn)
select {
case <-m.context.Done():
return trace.Wrap(m.Close())
default:
continue
}
}
if utils.IsUseOfClosedNetworkError(err) {
<-m.context.Done()
return nil
}
select {
case <-m.context.Done():
return nil
case <-time.After(5 * time.Second):
m.logger.LogAttrs(m.context, slog.LevelDebug, "Backoff on accept error", slog.Any("error", err))
}
}
}
// protocolListener returns a registered listener for Protocol proto
// and is safe for concurrent access.
func (m *Mux) protocolListener(proto Protocol) *Listener {
m.RLock()
defer m.RUnlock()
switch proto {
case ProtoTLS:
return m.tlsListener
case ProtoSSH:
return m.sshListener
case ProtoPostgres:
return m.dbListener
case ProtoHTTP:
return m.httpListener
}
return nil
}
// detectAndForward detects the protocol for conn and forwards to a
// registered protocol listener (SSH, TLS, DB). Connections for a
// protocol without a registered protocol listener are closed. This
// method is called as a goroutine by Serve for each connection.
func (m *Mux) detectAndForward(conn net.Conn) {
logger := m.logger.With(
"src_addr", logutils.StringerAttr(conn.RemoteAddr()),
"dst_addr", logutils.StringerAttr(conn.LocalAddr()),
)
if err := conn.SetDeadline(m.Clock.Now().Add(m.DetectTimeout)); err != nil {
logger.LogAttrs(m.context, slog.LevelWarn, "failed setting protocol detection deadline", slog.Any("error", err))
conn.Close()
return
}
var postDetect PostDetectFunc
if m.PreDetect != nil {
var err error
postDetect, err = m.PreDetect(conn)
if err != nil {
if !utils.IsOKNetworkError(err) {
logger.LogAttrs(m.context, slog.LevelWarn, "Failed to send early data",
slog.Any("error", err),
)
}
conn.Close()
return
}
}
connWrapper, err := m.detect(conn)
if err != nil {
if !errors.Is(trace.Unwrap(err), io.EOF) {
m.sampledLogger.LogAttrs(m.context, slog.LevelWarn, "failed to detect the connection type",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("error", err),
)
}
conn.Close()
return
}
if err := connWrapper.SetDeadline(time.Time{}); err != nil {
logger.WarnContext(m.context, "failed setting connection deadline", "error", err)
connWrapper.Close()
return
}
listener := m.protocolListener(connWrapper.protocol)
if listener == nil {
if connWrapper.protocol == ProtoHTTP {
logger.LogAttrs(m.context, slog.LevelDebug, "Detected an HTTP request - If this is for a health check, use an HTTPS request instead")
}
logger.LogAttrs(m.context, slog.LevelDebug, "Closing connection, listener is disabled",
slog.Any("protocol", logutils.StringerAttr(connWrapper.protocol)),
)
connWrapper.Close()
return
}
conn = connWrapper
if postDetect != nil {
conn = postDetect(connWrapper)
if conn == nil {
// the post detect hook hijacked the connection or had an error
return
}
}
listener.HandleConnection(m.context, conn)
}
// JWTPROXYSigner provides ability to created JWT for signed PROXY headers.
type JWTPROXYSigner interface {
SignPROXYJWT(p jwt.PROXYSignParams) (string, error)
}
func getTCPAddr(a net.Addr) net.TCPAddr {
if a == nil {
return net.TCPAddr{}
}
addr, ok := a.(*net.TCPAddr)
if ok { // Hot path
return *addr
}
parsedAddr := utils.FromAddr(a)
return net.TCPAddr{
IP: net.ParseIP(parsedAddr.Host()),
Port: parsedAddr.Port(-1),
}
}
func isDifferentTCPVersion(addr1, addr2 net.TCPAddr) bool {
return (addr1.IP.To4() != nil && addr2.IP.To4() == nil) || (addr2.IP.To4() != nil && addr1.IP.To4() == nil)
}
// hash an IPv6 into a class E IPv4
// https://developers.cloudflare.com/network/pseudo-ipv4/
func getPseudoIPV4(addr net.TCPAddr) (net.TCPAddr, error) {
hash := sha256.Sum256([]byte(addr.IP))
ip := hash[:4]
ip[0] |= classEPrefix
// don't assign the broadcast address
if slices.Equal(ip, []byte{255, 255, 255, 255}) {
ip[0] = 254
}
return net.TCPAddr{
IP: net.IP(ip),
Port: addr.Port,
}, nil
}
type signPROXYHeaderInput struct {
source net.Addr
destination net.Addr
allowDowngrade bool
clusterName string
signingCert []byte
signer JWTPROXYSigner
}
func signPROXYHeader(in signPROXYHeaderInput) ([]byte, error) {
originalSourceAddr := getTCPAddr(in.source)
sAddr := originalSourceAddr
dAddr := getTCPAddr(in.destination)
if sAddr.IP == nil || dAddr.IP == nil {
return nil, trace.Wrap(ErrBadIP, "source address: %s, destination address: %s", in.source, in.destination)
}
if sAddr.Port < 0 || dAddr.Port < 0 {
return nil, trace.BadParameter("could not parse port (source:%q, destination: %q)",
in.source.String(), in.destination.String())
}
if isDifferentTCPVersion(sAddr, dAddr) {
if !in.allowDowngrade {
return nil, trace.Wrap(ErrBadIP, "source address: %s, destination address: %s", in.source, in.destination)
}
// in a version mismatch, only the source address should be downgraded
if sAddr.IP.To4() != nil {
return nil, trace.Wrap(ErrDowngradeDst, "source address: %s, destination address: %s", in.source, in.destination)
}
var err error
if sAddr, err = getPseudoIPV4(sAddr); err != nil {
return nil, trace.Wrap(err)
}
// Mark original address, which will be returned as the RemoteAddr for Conns with a proxyLine configured, with port 0
// to prevent IP pinning. Pseudo IPv4 addresses are only made up of 31.5 bytes of sha256 hash which provides little
// defense against collisions
originalSourceAddr.Port = 0
}
signature, err := in.signer.SignPROXYJWT(jwt.PROXYSignParams{
SourceAddress: originalSourceAddr.String(),
DestinationAddress: dAddr.String(),
ClusterName: in.clusterName,
})
if err != nil {
return nil, trace.Wrap(err, "could not sign jwt token for PROXY line")
}
protocol := TCP4
if sAddr.IP.To4() == nil {
protocol = TCP6
}
pl := ProxyLine{
Protocol: protocol,
Source: sAddr,
Destination: dAddr,
}
var originalAddr *net.TCPAddr = nil
if !originalSourceAddr.IP.Equal(sAddr.IP) {
originalAddr = &originalSourceAddr
}
if err := pl.AddTeleportTLVs([]byte(signature), in.signingCert, originalAddr); err != nil {
return nil, trace.Wrap(err, "could not add signature to proxy line")
}
b, err := pl.Bytes()
if err != nil {
return nil, trace.Wrap(err, "could not get bytes from proxy line")
}
return b, nil
}
// errorSubstrings includes all the error substrings that can be returned by `Mux.detect`.
// These are used to deduplicate the errors returned by the multiplexer that occur
// when detecting the type of a new connection just established.
// This ensures that health checkers / malicious actors cannot pollute / overpower
// the logs with warnings when such connections are invalid or unknown to the multiplexer.
var errorSubstrings = []string{
failedToPeekConnectionError,
failedToDetectConnectionProtocolError,
externalProxyProtocolDisabledError,
duplicateSignedProxyLineError,
duplicateUnsignedProxyLineError,
invalidProxyLineError,
invalidProxyV2LineError,
invalidProxySignatureError,
unknownProtocolError,
missingProxyLineError,
unexpectedPROXYLineError,
unsignedPROXYLineAfterSignedError,
}
const (
// maxDetectionPasses sets maximum amount of passes to detect final protocol to account
// for 1 unsigned header, 1 signed header and the final protocol itself
maxDetectionPasses = 3
failedToPeekConnectionError = "failed to peek connection"
failedToDetectConnectionProtocolError = "failed to detect connection protocol"
externalProxyProtocolDisabledError = "external PROXY protocol support is disabled"
duplicateSignedProxyLineError = "duplicate signed PROXY line"
duplicateUnsignedProxyLineError = "duplicate unsigned PROXY line"
invalidProxyLineError = "invalid PROXY line"
invalidProxyV2LineError = "invalid PROXY v2 line"
invalidProxySignatureError = "could not verify PROXY signature for connection"
missingProxyLineError = `connection (%s -> %s) rejected: PROXY protocol required, but PROXY protocol line not received. Please verify your configuration.
Enable "proxy_protocol: on" only if Teleport is behind an L4 load balancer with PROXY protocol enabled.`
unknownProtocolError = "unknown protocol"
unexpectedPROXYLineError = `received unexpected PROXY protocol line. Connection will be allowed, but this is usually a result of misconfiguration -
if Teleport is running behind L4 load balancer with enabled PROXY protocol you should explicitly set config field "proxy_protocol" to "on".
See documentation for more details`
unsignedPROXYLineAfterSignedError = "received unsigned PROXY line after already receiving signed PROXY line"
)
// detect finds out a type of the connection and returns wrapper that support PROXY protocol
func (m *Mux) detect(conn net.Conn) (*Conn, error) {
reader := bufio.NewReader(conn)
// Before actual protocol traffic flows, we try to parse optional PROXY protocol headers,
// that can be injected by load balancers or our own proxies. There can be multiple PROXY
// headers. After they are parsed, last pass does the actual protocol detection itself.
// We allow only one unsigned PROXY header from external sources, if it's enabled, and one
// signed header from our own proxies, which take precedence.
var proxyLine *ProxyLine
unsignedPROXYLineReceived := false
for range maxDetectionPasses {
proto, err := detectProto(reader)
if err != nil {
return nil, trace.Wrap(err)
}
switch proto {
case ProtoProxy:
newPROXYLine, err := ReadProxyLine(reader)
if err != nil {
return nil, trace.Wrap(err, invalidProxyLineError)
}
if m.PROXYProtocolMode == PROXYProtocolOff {
return nil, trace.BadParameter("%s", externalProxyProtocolDisabledError)
}
if unsignedPROXYLineReceived {
// We allow only one unsigned PROXY line
return nil, trace.BadParameter("%s", duplicateUnsignedProxyLineError)
}
unsignedPROXYLineReceived = true
if m.PROXYProtocolMode == PROXYProtocolUnspecified && !m.SuppressUnexpectedPROXYWarning {
m.sampledLogger.LogAttrs(m.context, slog.LevelError, unexpectedPROXYLineError,
slog.Any("direct_src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("direct_dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("proxy_src_addr", logutils.StringerAttr(&newPROXYLine.Source)),
slog.Any("proxy_dst_addr", logutils.StringerAttr(&newPROXYLine.Destination)),
)
newPROXYLine.Source.Port = 0 // Mark connection, so if later IP pinning check is used on it we can reject it.
}
if proxyLine != nil && proxyLine.IsVerified {
// Unsigned PROXY line after signed one should not happen
return nil, trace.BadParameter("%s", unsignedPROXYLineAfterSignedError)
}
proxyLine = newPROXYLine
// repeat the cycle to detect the protocol
case ProtoProxyV2:
newPROXYLine, err := ReadProxyLineV2(reader)
if err != nil {
return nil, trace.Wrap(err, invalidProxyV2LineError)
}
if newPROXYLine == nil {
if unsignedPROXYLineReceived {
// We allow only one unsigned PROXY line
return nil, trace.BadParameter("%s", duplicateUnsignedProxyLineError)
}
unsignedPROXYLineReceived = true
continue // Skipping LOCAL command of PROXY protocol
}
// If proxyline is not signed, so we don't try to verify to avoid unnecessary load
if m.CertAuthorityGetter != nil && m.LocalClusterName != "" && newPROXYLine.IsSigned() {
err = newPROXYLine.VerifySignature(m.context, m.CertAuthorityGetter, m.LocalClusterName, m.Clock)
if errors.Is(err, ErrNoHostCA) {
m.logger.LogAttrs(m.context, slog.LevelWarn, "could not verify PROXY signature for connection, failed to get host CA",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
)
continue
}
if err != nil {
return nil, trace.Wrap(err, "%s %s -> %s", invalidProxySignatureError, conn.RemoteAddr(), conn.LocalAddr())
}
m.logger.LogAttrs(m.context, logutils.TraceLevel, "Successfully verified signed PROXYv2 header",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("client_src_addr", logutils.StringerAttr(&newPROXYLine.Source)),
)
}
// If proxy line is signed and successfully verified and there's no already signed proxy header,
// we accept, otherwise reject
if newPROXYLine.IsVerified {
if proxyLine != nil && proxyLine.IsVerified {
return nil, trace.BadParameter("%s", duplicateSignedProxyLineError)
}
proxyLine = newPROXYLine
continue
}
if m.CertAuthorityGetter != nil && newPROXYLine.IsSigned() && !newPROXYLine.IsVerified {
return nil, trace.BadParameter("could not verify PROXY line signature")
}
// This is unsigned proxy line, return error if external PROXY protocol is not enabled
if m.PROXYProtocolMode == PROXYProtocolOff {
return nil, trace.BadParameter("%s", externalProxyProtocolDisabledError)
}
if unsignedPROXYLineReceived {
// We allow only one unsigned PROXY line
return nil, trace.BadParameter("%s", duplicateUnsignedProxyLineError)
}
unsignedPROXYLineReceived = true
if m.PROXYProtocolMode == PROXYProtocolUnspecified && !m.SuppressUnexpectedPROXYWarning {
m.sampledLogger.LogAttrs(m.context, slog.LevelError, unexpectedPROXYLineError,
slog.Any("direct_src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("direct_dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("proxy_src_addr", logutils.StringerAttr(&newPROXYLine.Source)),
slog.Any("proxy_dst_addr", logutils.StringerAttr(&newPROXYLine.Destination)),
)
newPROXYLine.Source.Port = 0 // Mark connection, so if later IP pinning check is used on it we can reject it.
}
// Unsigned PROXY line after signed should not happen
if proxyLine != nil && proxyLine.IsVerified {
return nil, trace.BadParameter("%s", unsignedPROXYLineAfterSignedError)
}
proxyLine = newPROXYLine
// repeat the cycle to detect the protocol
case ProtoTLS, ProtoSSH, ProtoHTTP, ProtoPostgres:
if err := m.checkPROXYProtocolRequirement(conn, unsignedPROXYLineReceived); err != nil {
return nil, trace.Wrap(err)
}
return &Conn{
protocol: proto,
Conn: conn,
reader: reader,
proxyLine: proxyLine,
}, nil
}
}
// if code ended here after three attempts, something is wrong
return nil, trace.BadParameter("%s", unknownProtocolError)
}
// checkPROXYProtocolRequirement checks that if multiplexer is required to receive unsigned PROXY line
// that requirement is fulfilled, or exceptions apply - self connections and connections that are passed
// from upstream multiplexed listener (as it happens for alpn proxy).
func (m *Mux) checkPROXYProtocolRequirement(conn net.Conn, unsignedPROXYLineReceived bool) error {
if m.PROXYProtocolMode != PROXYProtocolOn {
return nil
}
// Proxy and other services might call itself directly, avoiding
// load balancer, so we shouldn't fail connections without PROXY headers for such cases.
selfConnection, err := m.isSelfConnection(conn)
if err != nil {
return trace.Wrap(err)
}
if !selfConnection && !isInternalConn(conn) && !unsignedPROXYLineReceived {
return trace.BadParameter(missingProxyLineError, conn.RemoteAddr().String(), conn.LocalAddr().String())
}
return nil
}
// isInternalConn determines if the connection is a multiplexer Conn.
// If the check is successful, it indicates that the connection was provided by another multiplexer listener,
// and that the unsigned PROXY protocol requirement has already been handled.
func isInternalConn(conn net.Conn) bool {
type netConn interface {
NetConn() net.Conn
}
for {
if _, ok := conn.(*Conn); ok {
return true
}
connGetter, ok := conn.(netConn)
if !ok {
return false
}
conn = connGetter.NetConn()
}
}
func (m *Mux) isSelfConnection(conn net.Conn) (bool, error) {
if m.IgnoreSelfConnections {
return false, nil
}
remoteHost, _, err := net.SplitHostPort(conn.RemoteAddr().String())
if err != nil {
return false, trace.Wrap(err)
}
localHost, _, err := net.SplitHostPort(conn.LocalAddr().String())
if err != nil {
return false, trace.Wrap(err)
}
return remoteHost == localHost, nil
}
// Protocol defines detected protocol type.
type Protocol int
const (
// ProtoUnknown is for unknown protocol
ProtoUnknown Protocol = iota
// ProtoTLS is TLS protocol
ProtoTLS
// ProtoSSH is SSH protocol
ProtoSSH
// ProtoProxy is a HAProxy proxy line protocol
ProtoProxy
// ProtoProxyV2 is a HAProxy binary protocol
ProtoProxyV2
// ProtoHTTP is HTTP protocol
ProtoHTTP
// ProtoPostgres is PostgreSQL wire protocol
ProtoPostgres
)
// protocolStrings defines strings for each Protocol.
var protocolStrings = map[Protocol]string{
ProtoUnknown: "Unknown",
ProtoTLS: "TLS",
ProtoSSH: "SSH",
ProtoProxy: "Proxy",
ProtoProxyV2: "ProxyV2",
ProtoHTTP: "HTTP",
ProtoPostgres: "Postgres",
}
// String returns the string representation of Protocol p.
// An empty string is returned when the protocol is not defined.
func (p Protocol) String() string {
return protocolStrings[p]
}
var (
proxyPrefix = []byte{'P', 'R', 'O', 'X', 'Y'}
ProxyV2Prefix = []byte{0x0D, 0x0A, 0x0D, 0x0A, 0x00, 0x0D, 0x0A, 0x51, 0x55, 0x49, 0x54, 0x0A}
sshPrefix = []byte{'S', 'S', 'H'}
tlsPrefix = []byte{0x16}
)
// This section defines Postgres wire protocol messages detected by Teleport:
//
// https://www.postgresql.org/docs/13/protocol-message-formats.html
var (
// postgresSSLRequest is always sent first by a Postgres client (e.g. psql)
// to check whether the server supports TLS.
postgresSSLRequest = []byte{0x0, 0x0, 0x0, 0x8, 0x4, 0xd2, 0x16, 0x2f}
// postgresCancelRequest is sent when a Postgres client requests
// cancellation of a long-running query.
//
// TODO(r0mant): It is currently unsupported because it is sent over a
// separate plain connection, but we're detecting it anyway so it at
// least appears in the logs as "unsupported" for debugging.
postgresCancelRequest = []byte{0x0, 0x0, 0x0, 0x10, 0x4, 0xd2, 0x16, 0x2e}
// postgresGSSEncRequest is sent first by a Postgres client
// to check whether the server supports GSS encryption.
// It is currently unsupported and our postgres engine will always respond 'N'
// for "not supported".
postgresGSSEncRequest = []byte{0x0, 0x0, 0x0, 0x8, 0x4, 0xd2, 0x16, 0x30}
)
var httpMethods = [...][]byte{
[]byte("GET"),
[]byte("POST"),
[]byte("PUT"),
[]byte("DELETE"),
[]byte("HEAD"),
[]byte("CONNECT"),
[]byte("OPTIONS"),
[]byte("TRACE"),
[]byte("PATCH"),
}
// isHTTP returns true if the first few bytes of the prefix indicate
// the use of an HTTP method.
func isHTTP(in []byte) bool {
for _, verb := range httpMethods {
if bytes.HasPrefix(in, verb) {
return true
}
}
return false
}
// detectProto tries to determine the network protocol used from the first
// few bytes of a connection.
func detectProto(r *bufio.Reader) (Protocol, error) {
// read the first 8 bytes without advancing the reader, some connections
// won't send more than 8 bytes at first
in, err := r.Peek(8)
if err != nil {
return ProtoUnknown, trace.Wrap(err, failedToPeekConnectionError)
}
switch {
case bytes.HasPrefix(in, proxyPrefix):
return ProtoProxy, nil
case bytes.HasPrefix(in, ProxyV2Prefix[:8]):
// if the first 8 bytes matches the first 8 bytes of the proxy
// protocol v2 magic bytes, read more of the connection so we can
// ensure all magic bytes match
in, err = r.Peek(len(ProxyV2Prefix))
if err != nil {
return ProtoUnknown, trace.Wrap(err, failedToPeekConnectionError)
}
if bytes.HasPrefix(in, ProxyV2Prefix) {
return ProtoProxyV2, nil
}
case bytes.HasPrefix(in, sshPrefix):
return ProtoSSH, nil
case bytes.HasPrefix(in, tlsPrefix):
return ProtoTLS, nil
case isHTTP(in):
return ProtoHTTP, nil
case bytes.HasPrefix(in, postgresSSLRequest),
bytes.HasPrefix(in, postgresCancelRequest),
bytes.HasPrefix(in, postgresGSSEncRequest):
return ProtoPostgres, nil
}
return ProtoUnknown, trace.BadParameter("%s, first few bytes were: %#v", failedToDetectConnectionProtocolError, in)
}
// PROXYHeaderSigner allows to sign PROXY headers for securely propagating original client IP information
type PROXYHeaderSigner interface {
SignPROXYHeader(source, destination net.Addr) ([]byte, error)
}
// PROXYSigner implements PROXYHeaderSigner to sign PROXY headers
type PROXYSigner struct {
getCertificate utils.GetCertificateFunc
clock clockwork.Clock
clusterName string
allowDowngrade bool
}
// NewPROXYSigner returns a new instance of PROXYSigner
func NewPROXYSigner(clusterName string, getCertificate utils.GetCertificateFunc, clock clockwork.Clock, allowDowngrade bool) (*PROXYSigner, error) {
return &PROXYSigner{
getCertificate: getCertificate,
clock: clock,
clusterName: clusterName,
allowDowngrade: allowDowngrade,
}, nil
}
// SignPROXYHeader creates a signed PROXY header with provided source and destination addresses
func (p *PROXYSigner) SignPROXYHeader(source, destination net.Addr) ([]byte, error) {
cert, err := p.getCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
if len(cert.Certificate) < 1 {
return nil, trace.Errorf("missing certificate for PROXY header signature")
}
if len(cert.Certificate) > 1 {
return nil, trace.Errorf("PROXY header signatures only support one certificate, got a chain of %v", len(cert.Certificate))
}
signingCert := cert.Certificate[0]
signer, ok := cert.PrivateKey.(crypto.Signer)
if !ok {
return nil, trace.Errorf("expected certificate private key to be a crypto.Signer, got %T", cert.PrivateKey)
}
jwtKey, err := jwt.New(&jwt.Config{
Clock: p.clock,
PrivateKey: signer,
ClusterName: p.clusterName,
})
if err != nil {
return nil, trace.Wrap(err)
}
proxyHeaderInput := signPROXYHeaderInput{
source: source,
destination: destination,
clusterName: p.clusterName,
signingCert: signingCert,
signer: jwtKey,
allowDowngrade: p.allowDowngrade,
}
header, err := signPROXYHeader(proxyHeaderInput)
if err != nil {
return nil, trace.Wrap(err)
}
if slog.Default().Enabled(context.Background(), logutils.TraceLevel) {
slog.LogAttrs(context.Background(), logutils.TraceLevel,
"Successfully signed PROXY header.",
slog.Any("src_addr", logutils.StringerAttr(source)),
slog.Any("dst_addr", logutils.StringerAttr(destination)),
slog.String("src_addr", p.clusterName),
)
}
return header, nil
}
/**
* Copyright 2013 Rackspace
*
* 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.
*
* Note: original copyright is preserved on purpose
*/
package multiplexer
import (
"bufio"
"bytes"
"context"
"crypto/x509"
"encoding/binary"
"encoding/hex"
"errors"
"fmt"
"io"
"math"
"net"
"slices"
"strconv"
"strings"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/jwt"
"github.com/gravitational/teleport/lib/tlsca"
)
// PP2Type is the PROXY protocol v2 TLV type
type PP2Type byte
type PP2TeleportSubtype PP2Type
const (
// TCP4 is TCP over IPv4
TCP4 = "TCP4"
// TCP6 is tCP over IPv6
TCP6 = "TCP6"
// Unknown is unsupported or unknown protocol
UNKNOWN = "UNKNOWN"
PP2TypeNOOP PP2Type = 0x04 // No-op used for padding
// Known custom types, spec allows to use 0xE0 - 0xEF for custom types
PP2TypeGCP PP2Type = 0xE0 // https://cloud.google.com/vpc/docs/configure-private-service-connect-producer
PP2TypeAWS PP2Type = 0xEA // https://docs.aws.amazon.com/elasticloadbalancing/latest/network/load-balancer-target-groups.html
PP2TypeAzure PP2Type = 0xEE // https://learn.microsoft.com/en-us/azure/private-link/private-link-service-overview
PP2TypeTeleport PP2Type = 0xE4 // Teleport own type for transferring our custom data such as connection metadata
PP2TeleportSubtypeSigningCert PP2TeleportSubtype = 0x01 // Certificate used to sign JWT
PP2TeleportSubtypeJWT PP2TeleportSubtype = 0x02 // JWT used to verify information sent in plain PROXY header
PP2TeleportSubtypeOriginalAddr PP2TeleportSubtype = 0x03 // Original IPv6 source address when downgrading to IPv4
)
var (
proxyCRLF = "\r\n"
proxySep = " "
// ErrTruncatedTLV is returned when there's no enough bytes to read full TLV
ErrTruncatedTLV = errors.New("TLV value was truncated")
// ErrNoSignature is returned when proxy line doesn't have full required data (JWT and cert) for verification
ErrNoSignature = errors.New("could not find signature data on the proxy line")
// ErrBadCACert is returned when a HostCA cert could not successfully be added to roots for signing certificate verification
ErrBadCACert = errors.New("could not add host CA to roots for verification")
// ErrIncorrectRole is returned when signing cert doesn't have required system role (Proxy)
ErrIncorrectRole = errors.New("could not find required system role on the signing certificate")
// ErrNonLocalCluster is returned when we received signed PROXY header, which signing certificate is from remote cluster.
ErrNonLocalCluster = errors.New("signing certificate is not signed by local cluster CA")
// ErrNoHostCA is returned when CAGetter could not get host CA, for example if auth server is not available
ErrNoHostCA = errors.New("could not get specified host CA to verify signed PROXY header")
// ErrInvalidPseudoIPv4 is returned when the proxy line source address is a pseudo IPv4 that does not match the signed IPv6
// included in the TLVs
ErrInvalidPseudoIPv4 = errors.New("mismatched pseudo IPv4 source and original IPv6 in proxy line")
)
// ProxyLine implements PROXY protocol version 1 and 2
// Spec: https://www.haproxy.org/download/1.8/doc/proxy-protocol.txt
// Original implementation here: https://github.com/racker/go-proxy-protocol
// TLV: https://github.com/pires/go-proxyproto
type ProxyLine struct {
Protocol string
Source net.TCPAddr
Destination net.TCPAddr
TLVs []TLV // PROXY protocol extensions
IsVerified bool
}
// TLV (Type-Length-Value) is an extension mechanism in PROXY protocol v2, see end of section 2.2
type TLV struct {
Type PP2Type
Value []byte
}
// String returns on-the wire string representation of the proxy line
func (p *ProxyLine) String() string {
return fmt.Sprintf("PROXY %s %s %s %d %d\r\n", p.Protocol, p.Source.IP.String(), p.Destination.IP.String(), p.Source.Port, p.Destination.Port)
}
// Bytes returns on-the wire bytes representation of proxy line conforming to the proxy v2 protocol
func (p *ProxyLine) Bytes() ([]byte, error) {
b := &bytes.Buffer{}
header := proxyV2Header{VersionCommand: (Version2 << 4) | ProxyCommand}
copy(header.Signature[:], ProxyV2Prefix)
var addr any
if p.Source.Port < 0 || p.Destination.Port < 0 ||
p.Source.Port > math.MaxUint16 || p.Destination.Port > math.MaxUint16 {
return nil, trace.BadParameter("source or destination port (%d,%d) is out of range 0-65535", p.Source.Port, p.Destination.Port)
}
switch p.Protocol {
case TCP4:
header.Protocol = ProtocolTCP4
addr4 := proxyV2Address4{
SourcePort: uint16(p.Source.Port),
DestinationPort: uint16(p.Destination.Port),
}
sourceIPv4 := p.Source.IP.To4()
if sourceIPv4 == nil {
return nil, trace.BadParameter("could not get source IPv4 address representation from %q", p.Source.IP.String())
}
copy(addr4.Source[:], sourceIPv4)
destIPv4 := p.Destination.IP.To4()
if destIPv4 == nil {
return nil, trace.BadParameter("could not get destination IPv4 address representation from %q", p.Destination.IP.String())
}
copy(addr4.Destination[:], destIPv4)
addr = addr4
case TCP6:
header.Protocol = ProtocolTCP6
addr6 := proxyV2Address6{
SourcePort: uint16(p.Source.Port),
DestinationPort: uint16(p.Destination.Port),
}
sourceIPv6 := p.Source.IP.To16()
if sourceIPv6 == nil {
return nil, trace.BadParameter("could not get source IPv6 address representation from %q", p.Source.IP.String())
}
copy(addr6.Source[:], sourceIPv6)
destIPv6 := p.Destination.IP.To16()
if destIPv6 == nil {
return nil, trace.BadParameter("could not get destination IPv6 address representation from %q", p.Destination.IP.String())
}
copy(addr6.Destination[:], destIPv6)
addr = addr6
default:
return nil, trace.BadParameter("unsupported protocol %q", p.Protocol)
}
tlvsBytes, err := MarshalTLVs(p.TLVs)
if err != nil {
return nil, trace.Errorf("could not marshal TLVs for the proxy line: %w", err)
}
if binary.Size(addr)+binary.Size(tlvsBytes) > math.MaxUint16 {
return nil, trace.LimitExceeded("size of PROXY payload is too large")
}
header.Length = uint16(binary.Size(addr) + binary.Size(tlvsBytes))
binary.Write(b, binary.BigEndian, header)
binary.Write(b, binary.BigEndian, addr)
binary.Write(b, binary.BigEndian, tlvsBytes)
return b.Bytes(), nil
}
// ReadProxyLine reads proxy line protocol from the reader
func ReadProxyLine(reader *bufio.Reader) (*ProxyLine, error) {
line, err := reader.ReadString('\n')
if err != nil {
return nil, trace.Wrap(err)
}
if !strings.HasSuffix(line, proxyCRLF) {
return nil, trace.BadParameter("expected CRLF in proxy protocol, got something else")
}
tokens := strings.Split(line[:len(line)-2], proxySep)
ret := ProxyLine{}
if len(tokens) < 6 {
return nil, trace.BadParameter("malformed PROXY line protocol string")
}
switch tokens[1] {
case TCP4:
ret.Protocol = TCP4
case TCP6:
ret.Protocol = TCP6
default:
ret.Protocol = UNKNOWN
}
sourceIP, err := parseIP(ret.Protocol, tokens[2])
if err != nil {
return nil, trace.Wrap(err)
}
destIP, err := parseIP(ret.Protocol, tokens[3])
if err != nil {
return nil, trace.Wrap(err)
}
sourcePort, err := parsePortNumber(tokens[4])
if err != nil {
return nil, trace.Wrap(err)
}
destPort, err := parsePortNumber(tokens[5])
if err != nil {
return nil, err
}
ret.Source = net.TCPAddr{IP: sourceIP, Port: sourcePort}
ret.Destination = net.TCPAddr{IP: destIP, Port: destPort}
return &ret, nil
}
func parsePortNumber(portString string) (int, error) {
port, err := strconv.Atoi(portString)
if err != nil {
return -1, trace.BadParameter("bad port %q: %v", portString, err)
}
if port < 0 || port > 65535 {
return -1, trace.BadParameter("port %q not in supported range [0...65535]", portString)
}
return port, nil
}
func parseIP(protocol string, addrString string) (net.IP, error) {
addr := net.ParseIP(addrString)
switch {
case len(addr) == 0:
return nil, trace.BadParameter("failed to parse address")
case addr.To4() != nil && protocol != TCP4:
return nil, trace.BadParameter("got IPV4 address %q for IPV6 proto %q", addr.String(), protocol)
case addr.To4() == nil && protocol != TCP6:
return nil, trace.BadParameter("got IPV6 address %v %q for IPV4 proto %q", len(addr), addr.String(), protocol)
}
return addr, nil
}
type proxyV2Header struct {
Signature [12]uint8
VersionCommand uint8
Protocol uint8
Length uint16
}
type proxyV2Address4 struct {
Source [4]uint8
Destination [4]uint8
SourcePort uint16
DestinationPort uint16
}
// proxyV2Address4Size is the size of a [proxyV2Address4]. Its correctness is
// enforced in proxyline_test.go to avoid having to import unsafe here.
const proxyV2Address4Size = 12
type proxyV2Address6 struct {
Source [16]uint8
Destination [16]uint8
SourcePort uint16
DestinationPort uint16
}
// proxyV2Address6Size is the size of a [proxyV2Address6]. Its correctness is
// enforced in proxyline_test.go to avoid having to import unsafe here.
const proxyV2Address6Size = 36
const (
Version2 = 2
ProxyCommand = 1
LocalCommand = 0
ProtocolTCP4 = 0x11
ProtocolTCP6 = 0x21
)
// ReadProxyLineV2 reads PROXY protocol v2 line from the reader
func ReadProxyLineV2(reader io.Reader) (*ProxyLine, error) {
var header proxyV2Header
var ret ProxyLine
if err := binary.Read(reader, binary.BigEndian, &header); err != nil {
return nil, trace.Wrap(err)
}
if !bytes.Equal(header.Signature[:], ProxyV2Prefix) {
return nil, trace.BadParameter("unrecognized signature %s", hex.EncodeToString(header.Signature[:]))
}
cmd, ver := header.VersionCommand&0xF, header.VersionCommand>>4
if ver != Version2 {
return nil, trace.BadParameter("unsupported version %d", ver)
}
if cmd == LocalCommand {
// LOCAL command, just skip address information and keep original addresses (no proxy line)
if header.Length > 0 {
_, err := io.CopyN(io.Discard, reader, int64(header.Length))
return nil, trace.Wrap(err)
}
return nil, nil
}
if cmd != ProxyCommand {
return nil, trace.BadParameter("unsupported command %d", cmd)
}
var size uint16
switch header.Protocol {
case ProtocolTCP4:
var addr proxyV2Address4
size = proxyV2Address4Size
if err := binary.Read(reader, binary.BigEndian, &addr); err != nil {
return nil, trace.Wrap(err)
}
ret.Protocol = TCP4
ret.Source = net.TCPAddr{IP: addr.Source[:], Port: int(addr.SourcePort)}
ret.Destination = net.TCPAddr{IP: addr.Destination[:], Port: int(addr.DestinationPort)}
case ProtocolTCP6:
var addr proxyV2Address6
size = proxyV2Address6Size
if err := binary.Read(reader, binary.BigEndian, &addr); err != nil {
return nil, trace.Wrap(err)
}
ret.Protocol = TCP6
ret.Source = net.TCPAddr{IP: addr.Source[:], Port: int(addr.SourcePort)}
ret.Destination = net.TCPAddr{IP: addr.Destination[:], Port: int(addr.DestinationPort)}
default:
return nil, trace.BadParameter("unsupported protocol %x", header.Protocol)
}
// If there are more bytes left it means we've got TLVs
if header.Length > size {
tlvsBytes := make([]byte, header.Length-size)
if _, err := io.ReadFull(reader, tlvsBytes); err != nil {
return nil, trace.Wrap(err)
}
tlvs, err := UnmarshalTLVs(tlvsBytes)
if err != nil {
return nil, trace.Wrap(err)
}
ret.TLVs = tlvs
}
return &ret, nil
}
// UnmarshalTLVs parses provided bytes slice into slice of TLVs
func UnmarshalTLVs(rawBytes []byte) ([]TLV, error) {
var tlvs []TLV
for len(rawBytes) > 0 {
if len(rawBytes) < 3 {
return nil, ErrTruncatedTLV
}
tlv := TLV{
Type: PP2Type(rawBytes[0]), // First byte is TLV type
}
// Next two bytes are TLV's value length
lenStart := 1
lenEnd := lenStart + 2
tlvLen := int(binary.BigEndian.Uint16(rawBytes[lenStart:lenEnd]))
rawBytes = rawBytes[3:] // Move by 3 bytes to skip TLV header
if tlvLen > len(rawBytes) {
return nil, ErrTruncatedTLV
}
// Ignore no-op padding
if tlv.Type != PP2TypeNOOP {
tlv.Value = make([]byte, tlvLen)
copy(tlv.Value, rawBytes[:tlvLen])
}
rawBytes = rawBytes[tlvLen:]
tlvs = append(tlvs, tlv)
}
return tlvs, nil
}
// MarshalTLVs marshals provided slice of TLVs into slice of bytes.
func MarshalTLVs(tlvs []TLV) ([]byte, error) {
var raw []byte
for _, tlv := range tlvs {
if len(tlv.Value) > math.MaxUint16 {
return nil, trace.LimitExceeded("can not marshal TLV with type %v, length %d exceeds the limit of 65kb", tlv.Type, len(tlv.Value))
}
var length [2]byte
binary.BigEndian.PutUint16(length[:], uint16(len(tlv.Value)))
raw = append(raw, byte(tlv.Type))
raw = append(raw, length[:]...)
raw = append(raw, tlv.Value...)
}
return raw, nil
}
// AddTeleportTLVs adds the provided signature, cert, and an optional original address to the proxy line,
// marshaling it into appropriate TLV structure.
func (p *ProxyLine) AddTeleportTLVs(signature, signingCert []byte, originalAddr *net.TCPAddr) error {
if len(signature) == 0 {
return trace.BadParameter("missing signature")
}
if len(signingCert) == 0 {
return trace.BadParameter("missing signing certificate")
}
teleportTLVs := []TLV{
{
Type: PP2Type(PP2TeleportSubtypeSigningCert),
Value: signingCert,
},
{
Type: PP2Type(PP2TeleportSubtypeJWT),
Value: signature,
},
}
if originalAddr != nil {
teleportTLVs = append(teleportTLVs, TLV{
Type: PP2Type(PP2TeleportSubtypeOriginalAddr),
Value: []byte(originalAddr.String()),
})
}
teleportTLVBytes, err := MarshalTLVs(teleportTLVs)
if err != nil {
return err
}
// If there's already signature among TLVs, we replace it
for i := range p.TLVs {
if p.TLVs[i].Type == PP2TypeTeleport {
p.TLVs[i].Value = teleportTLVBytes
return nil
}
}
// Otherwise we append it
p.TLVs = append(p.TLVs, TLV{Type: PP2TypeTeleport, Value: teleportTLVBytes})
return nil
}
// IsSigned returns true if proxy line's TLV contains signature.
// Does not take into account if signature is valid or not, just the presence of it.
func (p *ProxyLine) IsSigned() bool {
tlvs, _ := p.getTeleportTLVs()
return len(tlvs.token) > 0 || tlvs.proxyCert != nil
}
type teleportTLVs struct {
token string
proxyCert []byte
originalAddress *net.TCPAddr
}
// getTeleportTLVs returns custom teleport TLVs present in the ProxyLine, if any
func (p *ProxyLine) getTeleportTLVs() (teleportTLVs, error) {
var tlvs teleportTLVs
for _, tlv := range p.TLVs {
if tlv.Type == PP2TypeTeleport {
teleportSubTLVs, err := UnmarshalTLVs(tlv.Value)
if err != nil {
return tlvs, trace.Wrap(err)
}
for _, subTLV := range teleportSubTLVs {
switch PP2TeleportSubtype(subTLV.Type) {
case PP2TeleportSubtypeSigningCert:
tlvs.proxyCert = subTLV.Value
case PP2TeleportSubtypeJWT:
tlvs.token = string(subTLV.Value)
case PP2TeleportSubtypeOriginalAddr:
addr, err := net.ResolveTCPAddr("tcp", string(subTLV.Value))
if err != nil {
return tlvs, trace.Wrap(err)
}
// If the source address was marked with a port of 0 to prevent IP Pinning,
// then the original address should also have a port of 0.
if p.Source.Port == 0 {
addr.Port = 0
}
tlvs.originalAddress = addr
}
}
break
}
}
return tlvs, nil
}
// VerifySignature checks that signature contained in the proxy line is securely signed.
func (p *ProxyLine) VerifySignature(ctx context.Context, caGetter CertAuthorityGetter, localClusterName string, clock clockwork.Clock) error {
// If there's no TLVs it can't be verified
if len(p.TLVs) == 0 {
return trace.Wrap(ErrNoSignature)
}
tlvs, err := p.getTeleportTLVs()
if err != nil {
return trace.Wrap(err)
}
if len(tlvs.token) == 0 || tlvs.proxyCert == nil {
return trace.Wrap(ErrNoSignature)
}
signingCert, err := x509.ParseCertificate(tlvs.proxyCert)
if err != nil {
return trace.Wrap(err)
}
identity, err := tlsca.FromSubject(signingCert.Subject, signingCert.NotAfter)
if err != nil {
return trace.Wrap(err)
}
if identity.TeleportCluster != localClusterName {
return trace.Wrap(ErrNonLocalCluster, "signing certificate cluster name: %s, local cluster name: %s",
identity.TeleportCluster, localClusterName)
}
hostCA, err := caGetter(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: localClusterName,
}, false)
if err != nil {
return trace.Wrap(ErrNoHostCA, "CA cluster name: %s", localClusterName)
}
hostCACerts := getTLSCerts(hostCA)
roots := x509.NewCertPool()
for _, cert := range hostCACerts {
ok := roots.AppendCertsFromPEM(cert)
if !ok {
return trace.Wrap(ErrBadCACert)
}
}
// Make sure that transmitted proxy cert is signed by appropriate host CA
_, err = signingCert.Verify(x509.VerifyOptions{Roots: roots})
if err != nil {
return trace.Wrap(err)
}
foundRole := checkForSystemRole(identity, types.RoleProxy)
if !foundRole {
return trace.Wrap(ErrIncorrectRole)
}
// Check JWT using proxy cert's public key
jwtVerifier, err := jwt.New(&jwt.Config{
Clock: clock,
PublicKey: signingCert.PublicKey,
ClusterName: localClusterName,
})
if err != nil {
return trace.Wrap(err)
}
// Determine if a pseudo IPv4 was used and validate
sAddr := p.Source.String()
if tlvs.originalAddress != nil {
expectedPLSource, err := getPseudoIPV4(*tlvs.originalAddress)
if err != nil {
return trace.Wrap(err)
}
if !expectedPLSource.IP.Equal(p.Source.IP) {
return trace.Wrap(ErrInvalidPseudoIPv4)
}
sAddr = tlvs.originalAddress.String()
}
_, err = jwtVerifier.VerifyPROXY(jwt.PROXYVerifyParams{
ClusterName: localClusterName,
SourceAddress: sAddr,
DestinationAddress: p.Destination.String(),
RawToken: tlvs.token,
})
if err != nil {
return trace.Wrap(err)
}
p.IsVerified = true
return nil
}
// ResolveSource returns the source IP address associated with a ProxyLine. If the Source is a class E address
// then we need to return the IPv6 stored in the teleport TLVs instead.
func (p *ProxyLine) ResolveSource() net.Addr {
// check if class E address
if []byte(p.Source.IP)[0] < classEPrefix {
return &p.Source
}
if tlvs, err := p.getTeleportTLVs(); err == nil {
if tlvs.originalAddress != nil {
return tlvs.originalAddress
}
}
return &p.Source
}
func getTLSCerts(ca types.CertAuthority) [][]byte {
pairs := ca.GetTrustedTLSKeyPairs()
out := make([][]byte, len(pairs))
for i, pair := range pairs {
out[i] = slices.Clone(pair.Cert)
}
return out
}
func checkForSystemRole(identity *tlsca.Identity, roleToFind types.SystemRole) bool {
findRole := func(roles []string) bool {
for _, role := range roles {
if roleToFind == types.SystemRole(role) {
return true
}
}
return false
}
return findRole(identity.Groups) || findRole(identity.SystemRoles)
}
/*
* 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 multiplexer
import (
"context"
"io"
"log/slog"
"net"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/log/logtest"
)
// TestProxy is tcp passthrough proxy that sends a proxy-line when connecting
// to the target server.
type TestProxy struct {
listener net.Listener
target string
closeCh chan (struct{})
log *slog.Logger
v2 bool
}
// NewTestProxy creates a new test proxy that sends a proxy-line when
// proxying connections to the provided target address.
func NewTestProxy(target string, v2 bool) (*TestProxy, error) {
listener, err := net.Listen("tcp", "localhost:0")
if err != nil {
return nil, trace.Wrap(err)
}
return &TestProxy{
listener: listener,
target: target,
closeCh: make(chan struct{}),
log: logtest.With(teleport.ComponentKey, "test:proxy"),
v2: v2,
}, nil
}
// Address returns the proxy listen address.
func (p *TestProxy) Address() string {
return p.listener.Addr().String()
}
// Serve starts accepting client connections and proxying them to the target.
func (p *TestProxy) Serve() error {
for {
clientConn, err := p.listener.Accept()
if err != nil {
return trace.Wrap(err)
}
p.log.DebugContext(context.Background(), "Accepted connection", "remote_addr", logutils.StringerAttr(clientConn.RemoteAddr()))
go func() {
if err := p.handleConnection(clientConn); err != nil {
p.log.ErrorContext(context.Background(), "Failed to handle connection", "error", err)
}
}()
}
}
// handleConnection dials the target address, sends a proxy line to it and
// then starts proxying all traffic b/w client and target.
func (p *TestProxy) handleConnection(clientConn net.Conn) error {
serverConn, err := net.Dial("tcp", p.target)
if err != nil {
clientConn.Close()
return trace.Wrap(err)
}
defer serverConn.Close()
errCh := make(chan error, 2)
go func() { // Client -> server.
defer clientConn.Close()
defer serverConn.Close()
// Write proxy-line first and then start proxying from client.
err := p.sendProxyLine(clientConn, serverConn)
if err == nil {
_, err = io.Copy(serverConn, clientConn)
}
errCh <- trace.Wrap(err)
}()
go func() { // Server -> client.
defer clientConn.Close()
defer serverConn.Close()
_, err := io.Copy(clientConn, serverConn)
errCh <- trace.Wrap(err)
}()
var errs []error
for range 2 {
select {
case err := <-errCh:
if err != nil && !utils.IsOKNetworkError(err) {
errs = append(errs, err)
}
case <-p.closeCh:
p.log.DebugContext(context.Background(), "Closing")
return trace.NewAggregate(errs...)
}
}
return trace.NewAggregate(errs...)
}
// sendProxyLine sends proxy-line to the server.
func (p *TestProxy) sendProxyLine(clientConn, serverConn net.Conn) error {
clientAddr, err := utils.ParseAddr(clientConn.RemoteAddr().String())
if err != nil {
return trace.Wrap(err)
}
serverAddr, err := utils.ParseAddr(serverConn.RemoteAddr().String())
if err != nil {
return trace.Wrap(err)
}
proxyLine := &ProxyLine{
Protocol: TCP4,
Source: net.TCPAddr{IP: net.ParseIP(clientAddr.Host()), Port: clientAddr.Port(0)},
Destination: net.TCPAddr{IP: net.ParseIP(serverAddr.Host()), Port: serverAddr.Port(0)},
}
p.log.DebugContext(context.Background(), "Sending proxy line",
"proxy_line", proxyLine.String(),
"remote_addr", serverConn.RemoteAddr().String(),
)
if p.v2 {
b, bErr := proxyLine.Bytes()
if bErr != nil {
return trace.Wrap(err)
}
_, err = serverConn.Write(b)
} else {
_, err = serverConn.Write([]byte(proxyLine.String()))
}
if err != nil {
return trace.Wrap(err)
}
return nil
}
// Close closes the proxy listener.
func (p *TestProxy) Close() error {
close(p.closeCh)
return p.listener.Close()
}
/*
* 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 multiplexer
import (
"context"
"crypto/tls"
"errors"
"io"
"log/slog"
"net"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"golang.org/x/net/http2"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// TLSListenerConfig specifies listener configuration
type TLSListenerConfig struct {
// Listener is the listener returning *tls.Conn
// connections on Accept
Listener net.Listener
// ID is an identifier used for debugging purposes
ID string
// ReadDeadline is a connection read deadline during the TLS handshake (start
// of the connection). It is set to defaults.HandshakeReadDeadline if
// unspecified.
ReadDeadline time.Duration
// Clock is a clock to override in tests, set to real time clock
// by default
Clock clockwork.Clock
}
// CheckAndSetDefaults verifies configuration and sets defaults
func (c *TLSListenerConfig) CheckAndSetDefaults() error {
if c.Listener == nil {
return trace.BadParameter("missing parameter Listener")
}
if c.ReadDeadline == 0 {
c.ReadDeadline = defaults.HandshakeReadDeadline
}
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
return nil
}
// NewTLSListener returns a new TLS listener
func NewTLSListener(cfg TLSListenerConfig) (*TLSListener, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
context, cancel := context.WithCancel(context.TODO())
return &TLSListener{
log: slog.With(teleport.ComponentKey, teleport.Component("mxtls", cfg.ID)),
cfg: cfg,
http2Listener: newListener(context, cfg.Listener.Addr()),
httpListener: newListener(context, cfg.Listener.Addr()),
cancel: cancel,
context: context,
}, nil
}
// TLSListener wraps tls.Listener and detects negotiated protocol
// (assuming it's either http/1.1 or http/2)
// and forwards the appropriate responses to either HTTP/1.1 or HTTP/2
// listeners
type TLSListener struct {
log *slog.Logger
cfg TLSListenerConfig
http2Listener *Listener
httpListener *Listener
cancel context.CancelFunc
context context.Context
}
// HTTP2 returns HTTP2 listener
func (l *TLSListener) HTTP2() net.Listener {
return l.http2Listener
}
// HTTP returns HTTP listener
func (l *TLSListener) HTTP() net.Listener {
return l.httpListener
}
// Serve accepts and forwards tls.Conn connections
func (l *TLSListener) Serve() error {
for {
conn, err := l.cfg.Listener.Accept()
if err == nil {
tlsConn, ok := conn.(*tls.Conn)
if !ok {
conn.Close()
l.log.LogAttrs(l.context, slog.LevelError, "Received a non-TLS connection",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("conn_type", logutils.TypeAttr(conn)),
)
continue
}
go l.detectAndForward(tlsConn)
continue
}
if utils.IsUseOfClosedNetworkError(err) {
<-l.context.Done()
return nil
}
select {
case <-l.context.Done():
return nil
case <-time.After(5 * time.Second):
}
}
}
func (l *TLSListener) detectAndForward(conn *tls.Conn) {
err := conn.SetReadDeadline(l.cfg.Clock.Now().Add(l.cfg.ReadDeadline))
if err != nil {
l.log.LogAttrs(l.context, slog.LevelDebug, "Failed to set connection deadline",
slog.Any("error", err),
)
conn.Close()
return
}
start := l.cfg.Clock.Now()
if err := conn.HandshakeContext(l.context); err != nil {
if !errors.Is(trace.Unwrap(err), io.EOF) {
l.log.LogAttrs(l.context, slog.LevelWarn, "Handshake failed",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("error", err),
)
}
conn.Close()
return
}
// Log warning if TLS handshake takes more than one second to help debug
// latency issues.
if elapsed := time.Since(start); elapsed > 1*time.Second {
l.log.LogAttrs(l.context, slog.LevelWarn, "Slow TLS handshake",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Duration("handshake_duration", time.Since(start)),
)
}
err = conn.SetReadDeadline(time.Time{})
if err != nil {
l.log.WarnContext(l.context, "Failed to reset read deadline", "error", err)
conn.Close()
return
}
switch conn.ConnectionState().NegotiatedProtocol {
case http2.NextProtoTLS:
l.http2Listener.HandleConnection(l.context, conn)
case teleport.HTTPNextProtoTLS, "":
l.httpListener.HandleConnection(l.context, conn)
default:
conn.Close()
l.log.LogAttrs(l.context, slog.LevelError, "rejecting connection with unsupported protocol",
slog.Any("error", err),
slog.String("protocol", conn.ConnectionState().NegotiatedProtocol),
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
)
}
}
// Close closes the listener.
// Any blocked Accept operations will be unblocked and return errors.
func (l *TLSListener) Close() error {
defer l.cancel()
return l.cfg.Listener.Close()
}
// Addr returns the listener's network address.
func (l *TLSListener) Addr() net.Addr {
return l.cfg.Listener.Addr()
}
/*
* 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 multiplexer
import (
"context"
"crypto/tls"
"errors"
"io"
"log/slog"
"net"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/defaults"
dbcommon "github.com/gravitational/teleport/lib/srv/db/dbutils"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// WebListenerConfig is the web listener configuration.
type WebListenerConfig struct {
// Listener is the listener that accepts tls connections.
Listener net.Listener
// ReadDeadline is a connection read deadline during the TLS handshake.
ReadDeadline time.Duration
// Clock is a clock to override in tests.
Clock clockwork.Clock
}
// CheckAndSetDefaults verifies configuration and sets defaults.
func (c *WebListenerConfig) CheckAndSetDefaults() error {
if c.Listener == nil {
return trace.BadParameter("missing parameter Listener")
}
if c.ReadDeadline == 0 {
c.ReadDeadline = defaults.HandshakeReadDeadline
}
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
return nil
}
// NewWebListener returns a new web listener.
func NewWebListener(cfg WebListenerConfig) (*WebListener, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
context, cancel := context.WithCancel(context.Background())
return &WebListener{
log: slog.With(teleport.ComponentKey, "mxweb"),
cfg: cfg,
webListener: newListener(context, cfg.Listener.Addr()),
dbListener: newListener(context, cfg.Listener.Addr()),
cancel: cancel,
context: context,
}, nil
}
// WebListener multiplexes tls connections between web and database listeners
// based on the client certificate.
type WebListener struct {
log *slog.Logger
cfg WebListenerConfig
webListener *Listener
dbListener *Listener
cancel context.CancelFunc
context context.Context
}
// Web returns web listener.
func (l *WebListener) Web() net.Listener {
return l.webListener
}
// DB returns database access listener.
func (l *WebListener) DB() net.Listener {
return l.dbListener
}
// Serve starts accepting and forwarding tls connections to appropriate listeners.
func (l *WebListener) Serve() error {
for {
conn, err := l.cfg.Listener.Accept()
if err != nil {
if utils.IsUseOfClosedNetworkError(err) {
<-l.context.Done()
return trace.Wrap(err, "listener is closed")
}
select {
case <-l.context.Done():
return trace.Wrap(net.ErrClosed, "listener is closed")
case <-time.After(5 * time.Second):
l.log.LogAttrs(l.context, slog.LevelWarn, "Backoff on accept error",
slog.Any("error", err),
)
}
continue
}
tlsConn, ok := conn.(*tls.Conn)
if !ok {
l.log.LogAttrs(l.context, slog.LevelError, "Received a non-TLS connection",
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
slog.Any("conn_type", logutils.TypeAttr(conn)),
)
conn.Close()
continue
}
go l.detectAndForward(tlsConn)
}
}
func (l *WebListener) detectAndForward(conn *tls.Conn) {
err := conn.SetReadDeadline(l.cfg.Clock.Now().Add(l.cfg.ReadDeadline))
if err != nil {
l.log.LogAttrs(l.context, slog.LevelWarn, "Failed to set connection read deadline",
slog.Any("error", err),
)
conn.Close()
return
}
if err := conn.HandshakeContext(l.context); err != nil {
if !errors.Is(trace.Unwrap(err), io.EOF) {
l.log.LogAttrs(l.context, slog.LevelWarn, "Handshake failed",
slog.Any("error", err),
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
)
}
conn.Close()
return
}
err = conn.SetReadDeadline(time.Time{})
if err != nil {
l.log.WarnContext(l.context, "Failed to reset connection read deadline", "error", err)
conn.Close()
return
}
// Inspect the client certificate (if there's any) and forward the
// connection either to database access listener if identity encoded
// in the cert indicates this is a database connection, or to a regular
// tls listener.
isDatabaseConnection, err := dbcommon.IsDatabaseConnection(conn.ConnectionState())
if err != nil {
l.log.LogAttrs(l.context, slog.LevelDebug, "Failed to check if connection is database connection",
slog.Any("error", err),
slog.Any("src_addr", logutils.StringerAttr(conn.RemoteAddr())),
slog.Any("dst_addr", logutils.StringerAttr(conn.LocalAddr())),
)
}
if isDatabaseConnection {
l.dbListener.HandleConnection(l.context, conn)
return
}
l.webListener.HandleConnection(l.context, conn)
}
// Close closes the listener.
//
// Any blocked Accept operations will be unblocked and return errors.
func (l *WebListener) Close() error {
defer l.cancel()
return l.cfg.Listener.Close()
}
// Addr returns the listener's network address.
func (l *WebListener) Addr() net.Addr {
return l.cfg.Listener.Addr()
}
/*
* 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 multiplexer
import (
"bufio"
"context"
"net"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/utils"
)
// Conn is a connection wrapper that supports
// communicating remote address from proxy protocol
// and replays first several bytes read during
// protocol detection
type Conn struct {
net.Conn
protocol Protocol
proxyLine *ProxyLine
reader *bufio.Reader
}
// NewConn returns a net.Conn wrapper that supports peeking into the connection.
func NewConn(conn net.Conn) *Conn {
return &Conn{
Conn: conn,
reader: bufio.NewReader(conn),
}
}
// NetConn returns the underlying net.Conn.
func (c *Conn) NetConn() net.Conn {
return c.Conn
}
// Read reads from connection
func (c *Conn) Read(p []byte) (int, error) {
return c.reader.Read(p)
}
// Peek is [*bufio.Reader.Peek].
func (c *Conn) Peek(n int) ([]byte, error) {
return c.reader.Peek(n)
}
// Discard is [*bufio.Reader.Discard].
func (c *Conn) Discard(n int) (discarded int, err error) {
return c.reader.Discard(n)
}
// ReadByte is [*bufio.Reader.ReadByte].
func (c *Conn) ReadByte() (byte, error) {
return c.reader.ReadByte()
}
// LocalAddr returns local address of the connection
func (c *Conn) LocalAddr() net.Addr {
if c.proxyLine != nil {
return &c.proxyLine.Destination
}
return c.Conn.LocalAddr()
}
// RemoteAddr returns remote address of the connection
func (c *Conn) RemoteAddr() net.Addr {
if c.proxyLine != nil {
return c.proxyLine.ResolveSource()
}
return c.Conn.RemoteAddr()
}
// Protocol returns the detected connection protocol
func (c *Conn) Protocol() Protocol {
return c.protocol
}
// Detect detects the connection protocol by peeking into the first few bytes.
func (c *Conn) Detect() (Protocol, error) {
proto, err := detectProto(c.reader)
if err != nil && !trace.IsBadParameter(err) {
return ProtoUnknown, trace.Wrap(err)
}
return proto, nil
}
// ReadProxyLine reads proxy-line from the connection.
func (c *Conn) ReadProxyLine() (*ProxyLine, error) {
var proxyLine *ProxyLine
protocol, err := c.Detect()
if err != nil {
return nil, trace.Wrap(err)
}
if protocol == ProtoProxyV2 {
proxyLine, err = ReadProxyLineV2(c.reader)
} else {
proxyLine, err = ReadProxyLine(c.reader)
}
if err != nil {
return nil, trace.Wrap(err)
}
c.proxyLine = proxyLine
return proxyLine, nil
}
// returns a Listener that pretends to be listening on addr, closed whenever the
// parent context is done.
func newListener(parent context.Context, addr net.Addr) *Listener {
context, cancel := context.WithCancel(parent)
return &Listener{
addr: addr,
connC: make(chan net.Conn),
cancel: cancel,
context: context,
}
}
// Listener is a listener that receives
// connections from multiplexer based on the connection type
type Listener struct {
addr net.Addr
connC chan net.Conn
cancel context.CancelFunc
context context.Context
}
// Addr returns listener addr, the address of multiplexer listener
func (l *Listener) Addr() net.Addr {
return l.addr
}
// Accept accepts connections from parent multiplexer listener
func (l *Listener) Accept() (net.Conn, error) {
select {
case <-l.context.Done():
return nil, trace.ConnectionProblem(net.ErrClosed, "listener is closed")
case conn := <-l.connC:
return conn, nil
}
}
// HandleConnection injects the connection into the Listener, blocking until the
// context expires, the connection is accepted or the Listener is closed.
func (l *Listener) HandleConnection(ctx context.Context, conn net.Conn) {
select {
case <-ctx.Done():
conn.Close()
case <-l.context.Done():
conn.Close()
case l.connC <- conn:
}
}
// Close closes the listener.
func (l *Listener) Close() error {
l.cancel()
return nil
}
// PROXYEnabledListener wraps provided listener and can receive and apply PROXY headers and then pass connection up the chain.
type PROXYEnabledListener struct {
cfg Config
mux *Mux
net.Listener
}
// NewPROXYEnabledListener creates news instance of PROXYEnabledListener
func NewPROXYEnabledListener(cfg Config) (*PROXYEnabledListener, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
mux, err := New(cfg) // Creating Mux to leverage protocol detection with PROXY headers
if err != nil {
return nil, trace.Wrap(err)
}
muxListener := mux.SSH()
go func() {
if err := mux.Serve(); err != nil && !utils.IsOKNetworkError(err) {
mux.logger.ErrorContext(cfg.Context, "Mux encountered err serving", "error", err)
}
}()
pl := &PROXYEnabledListener{
cfg: cfg,
mux: mux,
Listener: muxListener,
}
return pl, nil
}
func (p *PROXYEnabledListener) Close() error {
return trace.Wrap(p.mux.Close())
}
// Accept gets connection from the wrapped listener and detects whether we receive PROXY headers on it,
// after first non PROXY protocol detected it returns connection with PROXY addresses applied to it.
func (p *PROXYEnabledListener) Accept() (net.Conn, error) {
conn, err := p.Listener.Accept()
return conn, trace.Wrap(err)
}
// Teleport
// Copyright (C) 2026 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 scopes
import (
"encoding/hex"
"slices"
"strings"
"github.com/gravitational/trace"
)
// ResourceCursorPrefix prefixes cursors for scoped resources in a
// logical resource stream.
//
// The prefix starts with '~', which is not allowed in backend-safe resource
// names and sorts after all backend-safe name bytes, preserving historical
// name-only cursors for unscoped resources while ordering scoped resources
// after unscoped resources.
const ResourceCursorPrefix = "~scoped/"
// ResourceCursorScopedStart returns the first cursor in the scoped portion of
// the logical resource stream.
func ResourceCursorScopedStart() string {
return ResourceCursorPrefix
}
// IsScopedResourceCursor returns true if cursor is in the scoped portion of the
// logical resource stream.
func IsScopedResourceCursor(cursor string) bool {
return strings.HasPrefix(cursor, ResourceCursorPrefix)
}
// MakeResourceCursor returns the cursor for a scoped or unscoped resource in a
// logical, lexicographically ordered resource stream.
//
// Resource cursors are intended for pagination tokens, range bounds, and
// in-memory cache indexes. They are not backend storage keys and must not be
// used to construct backend keys.
//
// Unscoped resource cursors preserve the historical name-only format:
//
// <name>
//
// Scoped resource cursors use a synthetic prefix that cannot appear in
// backend-safe resource names and sorts after all backend-safe name bytes:
//
// ~scoped/<encoded-scope>/<name>
//
// The scope component is encoded with [EncodeForKey] so that scoped cursors
// preserve scope ordering and can safely use '/' as the cursor separator.
//
// MakeResourceCursor is infallible so that it can back key derivation with no
// error path (in-memory cache indexes, pagination cursors). A scope that
// cannot be encoded (which is only possible for invalid stored data) yields a
// degraded cursor that is deterministic, unique per scope and name, sorts after
// all valid cursors, and fails [ParseResourceCursor].
func MakeResourceCursor(scope, name string) string {
return MakeNestedResourceCursor(QualifiedName{
Scope: scope,
Name: name,
})
}
// MakeNestedResourceCursor returns the cursor for a scoped or unscoped nested
// resource in a logical, lexicographically ordered range of nested resources.
// The motivating example is scoped access list members, where members are
// keyed under their parent list.
//
// Resource cursors are intended for pagination tokens, range bounds, and
// in-memory cache indexes. They are not backend storage keys and must not be
// used to construct backend keys.
//
// If all provided scopes are empty, all names will simply be joined with the
// separator, to maintain compatibility with existing unscoped resource cursors
// and avoid wastefully encoding multiple empty scopes:
//
// <root-name>[/<descendent-name>]...
//
// Scoped resource cursors use a synthetic prefix that cannot appear in
// backend-safe resource names and sorts after all backend-safe name bytes.
//
// ~scoped/<encoded-root-scope>/<root-name>[/<encoded-descendent-scope>/<descendent-name>]...
//
// Each scope component is encoded with [EncodeForKey] so that scoped cursors
// preserve scope ordering and can safely use '/' as the cursor separator.
//
// MakeNestedResourceCursor is infallible so that it can back key derivation with no
// error path (in-memory cache indexes, pagination cursors). A scope that
// cannot be encoded (which is only possible for invalid stored data) yields a
// degraded cursor that is deterministic, unique per scope and name, sorts after
// all valid cursors, and fails [ParseResourceCursor].
func MakeNestedResourceCursor(root QualifiedName, descendents ...QualifiedName) string {
hasNonEmptyScope := root.Scope != "" || slices.ContainsFunc(descendents, func(descendent QualifiedName) bool {
return descendent.Scope != ""
})
if !hasNonEmptyScope {
var sb strings.Builder
sb.WriteString(root.Name)
for _, descendent := range descendents {
sb.WriteString(separator)
sb.WriteString(descendent.Name)
}
return sb.String()
}
var sb strings.Builder
sb.WriteString(ResourceCursorPrefix)
sb.WriteString(EncodeForResourceCursor(root.Scope))
sb.WriteString(separator)
sb.WriteString(root.Name)
for _, descendent := range descendents {
sb.WriteString(separator)
sb.WriteString(EncodeForResourceCursor(descendent.Scope))
sb.WriteString(separator)
sb.WriteString(descendent.Name)
}
return sb.String()
}
// EncodeForResourceCursor is infallible so that it can back key derivation with no
// error path (in-memory cache indexes, pagination cursors). A scope that
// cannot be encoded (which is only possible for invalid stored data) yields a
// degraded cursor that is deterministic, unique per scope and name, sorts after
// all valid cursors, and fails [ParseResourceCursor].
func EncodeForResourceCursor(scope string) string {
encoded, err := EncodeForKey(scope)
if err != nil {
// '~' cannot appear in a valid scope encoding (which starts with '+'),
// so degraded cursors never collide with valid cursors and sort after
// them. The scope is hex-encoded so the cursor's scope component is a
// single path segment regardless of the scope's contents.
return "~invalid+" + hex.EncodeToString([]byte(scope))
}
return encoded
}
// MakeResourceCursorWithHost returns the cursor for a scoped or unscoped
// host-keyed resource — one keyed by (host ID, name), such as an app server —
// in a logical, lexicographically ordered resource stream:
//
// scoped: ~scoped/<encoded-scope>/<host-id>/<name>
//
// See [MakeResourceCursor] for cursor semantics. Host cursors are not
// parseable by [ParseResourceCursor], which rejects composite names.
func MakeResourceCursorWithHost(scope, hostID, name string) string {
return MakeResourceCursor(scope, hostID+separator+name)
}
// ParseResourceCursor parses a resource cursor produced by [MakeResourceCursor]
// into its scope and name components.
//
// Unscoped cursors are interpreted as historical name-only cursors. Scoped
// cursors must use the scoped cursor format:
//
// ~scoped/<encoded-scope>/<name>
func ParseResourceCursor(cursor string) (QualifiedName, error) {
encodedScopeAndName, ok := strings.CutPrefix(cursor, ResourceCursorPrefix)
if !ok {
return QualifiedName{Name: cursor}, nil
}
encodedScope, name, ok := strings.Cut(encodedScopeAndName, separator)
if !ok {
return QualifiedName{}, trace.BadParameter("scoped resource cursor %q missing name separator", cursor)
}
if encodedScope == "" {
return QualifiedName{}, trace.BadParameter("scoped resource cursor %q has empty encoded scope", cursor)
}
if name == "" {
return QualifiedName{}, trace.BadParameter("scoped resource cursor %q has empty name", cursor)
}
if strings.Contains(name, separator) {
return QualifiedName{}, trace.BadParameter("scoped resource cursor %q has invalid name", cursor)
}
scope, err := DecodeFromKey(encodedScope)
if err != nil {
return QualifiedName{}, trace.Wrap(err)
}
return QualifiedName{Scope: scope, Name: name}, nil
}
// Teleport
// Copyright (C) 2026 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 scopes
import (
"encoding/hex"
"github.com/gravitational/trace"
"rsc.io/ordered"
)
// The scope key encoding produces a single opaque, order-preserving backend key
// segment for a scope. It is used to namespace scoped resources in the backend,
// e.g. as the <encoded_scope> component of a key like:
//
// /scoped/<kind>/<encoded_scope>/<name>
//
// The encoding is built in two phases. First, the scope is encoded as an
// [ordered] sequence consisting of a leading discriminant that distinguishes
// scoped from unscoped values, followed by one string element per scope segment:
//
// unscoped ("") -> ordered.Encode(unscopedDisc)
// root ("/") -> ordered.Encode(scopedDisc)
// "/a" -> ordered.Encode(scopedDisc, "a")
// "/a/b" -> ordered.Encode(scopedDisc, "a", "b")
//
// This encoding was chosen to satisfy several properties simultaneously:
//
// - Order-preserving: a plain byte-sort of the encoded values reproduces
// the same sort order as the [Sort] function.
//
// - Scope prefixing: An encoded scope S is a string prefix of any encoded child scope,
// but is *not* a prefix of a scope that is a sibling of S with the same leading segment
// characters (e.g. EncodeForKey("/staging") is a prefix of EncodeForKey("/staging/west")
// but not of EncodeForKey("/stagingwest")).
//
// - Exact prefixing in backend: Appending the backend separator to the encoded scope allows
// backend range queries to retrieve *exactly* the set of keys with that exact scope prefix,
// without ambiguity.
//
// - Forward-compatible: The encoding scheme can theoretically handle any future extension to
// allowed scope characters.
const (
// scopeKeyUnscopedDisc is the leading discriminant for the encoding of an
// unscoped value (i.e. EncodeForKey("")). It sorts before
// scopeKeyScopedDisc so that unscoped values sort before all scoped values.
scopeKeyUnscopedDisc = 0
// scopeKeyScopedDisc is the leading discriminant for the encoding of any
// scoped value (including the root scope "/").
scopeKeyScopedDisc = 1
)
// EncodeForKey encodes a scope so that it will be valid for use in a single
// backend key segment, preserving sort order. If given an empty string, it
// will return a non-empty encoding that sorts before all valid encoded
// scopes.
func EncodeForKey(scope string) (string, error) {
if scope == "" {
return hex.EncodeToString(ordered.Encode(scopeKeyUnscopedDisc)), nil
}
if err := WeakValidate(scope); err != nil {
return "", trace.Wrap(err)
}
raw := ordered.Encode(scopeKeyScopedDisc)
for segment := range DescendingSegments(scope) {
raw = ordered.Append(raw, segment)
}
return hex.EncodeToString(raw), nil
}
// DecodeFromKey decodes a scope encoded by [EncodeForKey].
func DecodeFromKey(encoded string) (string, error) {
raw, err := hex.DecodeString(encoded)
if err != nil {
return "", trace.BadParameter("invalid encoded scope %q: %v", encoded, err)
}
if len(raw) == 0 {
return "", trace.BadParameter("invalid empty encoded scope")
}
var disc int
rest, err := ordered.DecodePrefix(raw, &disc)
if err != nil {
return "", trace.BadParameter("malformed encoded scope %q: %v", encoded, err)
}
switch disc {
case scopeKeyUnscopedDisc:
if len(rest) != 0 {
return "", trace.BadParameter("malformed unscoped encoding %q: unexpected trailing data", encoded)
}
return "", nil
case scopeKeyScopedDisc:
var segments []string
for len(rest) > 0 {
var segment string
if rest, err = ordered.DecodePrefix(rest, &segment); err != nil {
return "", trace.BadParameter("malformed encoded scope %q: %v", encoded, err)
}
segments = append(segments, segment)
}
decoded := Join(segments...)
if err := WeakValidate(decoded); err != nil {
return "", trace.Wrap(err)
}
return decoded, nil
default:
return "", trace.BadParameter("invalid scope discriminant %d in encoded scope %q", disc, encoded)
}
}
/*
* Teleport
* Copyright (C) 2025 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 scopes
import (
"os"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/proto"
apiutils "github.com/gravitational/teleport/api/utils"
)
const (
// featureVarName is the name of the unstable scopes feature flag.
featureVarName = "TELEPORT_UNSTABLE_SCOPES"
// agentPinVarName is the name of the unstable agent scope pin feature flag.
agentPinVarName = "TELEPORT_UNSTABLE_AGENT_SCOPE_PIN"
)
// Features describes which scopes-related functionality is enabled.
type Features struct {
// Enabled indicates whether the base scopes feature is enabled.
Enabled bool
// AgentPinEnabled checks if the agent scope pin feature is enabled.
AgentPinEnabled bool
}
// AssertEnabled returns an error if the base scopes feature is disabled.
func (f Features) AssertEnabled() error {
if !f.Enabled {
return trace.Errorf("scoping features are not enabled, set " + featureVarName + "=yes to enable scoping features (caution: not ready for production use)")
}
return nil
}
// FeaturesFromEnv builds Features from scopes-related environment variables.
func FeaturesFromEnv() Features {
var f Features
enabled, err := apiutils.ParseBool(os.Getenv(featureVarName))
f.Enabled = enabled && err == nil
agentPinEnabled, err := apiutils.ParseBool(os.Getenv(agentPinVarName))
f.AgentPinEnabled = agentPinEnabled && err == nil
return f
}
// ScopesStatusToString returns a user friendly status message based on [proto.ScopesStatus].
func ScopesStatusToString(s proto.ScopesStatus) string {
switch s {
case proto.ScopesStatus_SCOPES_STATUS_ENABLED:
return "enabled"
case proto.ScopesStatus_SCOPES_STATUS_DISABLED:
return "disabled"
default:
return "unknown"
}
}
/*
* Teleport
* Copyright (C) 2026 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 scopes
import (
"github.com/gravitational/trace"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
)
// IsMatchAll reports whether the given filter is a wildcard match that selects all resources.
func IsMatchAll(filter *scopesv1.Filter) bool {
// an unspecified filter and MODE_ALL are both treated as a wildcard match at the matching layer. see
// [MatchScope] for discussion of why these are equivalent here despite being authorized differently
// at the API/authz layer.
mode := filter.GetMode()
return mode == scopesv1.Mode_MODE_UNSPECIFIED || mode == scopesv1.Mode_MODE_ALL
}
// MatchScope reports whether a resource at the given scope matches the supplied filter. This is the
// authoritative scope-matching logic used by caches/backends to decide which resources a filter selects.
//
// The empty string ("") is used to represent an unscoped resource and is treated as orthogonal to all
// scoped values (it is *not* equivalent to the root scope). As a result, the relationship-based modes
// never match unscoped resources, and MODE_UNSCOPED matches only unscoped resources.
//
// Matching logic here treats empty/unspecified filters and MODE_ALL as equivalent wildcard matchers. However,
// the API/authz defaults an empty/unspecified filter to being one of UNSCOPED or EXACT depending on the
// caller's identity. This helps ensure that outdated or un-explicit calls result in safe/conservative defaults.
//
// This function does not validate its inputs. Use [ValidateFilter] to verify that a filter is well-formed
// (e.g. that the scope and mode are mutually consistent) before relying on it.
func MatchScope(filter *scopesv1.Filter, resourceScope string) bool {
if IsMatchAll(filter) {
return true
}
mode := filter.GetMode()
if mode == scopesv1.Mode_MODE_UNSCOPED {
// unscoped resources only.
return resourceScope == ""
}
// all remaining modes select resources by the relationship of the resource scope to the filter scope.
rel := Compare(filter.GetScope(), resourceScope)
switch mode {
case scopesv1.Mode_MODE_EXACT:
return rel == Equivalent
case scopesv1.Mode_MODE_DESCENDANTS:
return rel == Equivalent || rel == Descendant
case scopesv1.Mode_MODE_ANCESTORS, scopesv1.Mode_MODE_POLICIES_APPLICABLE_TO_SCOPE: //nolint:staticcheck // SA1019. Deprecated mode retained as equivalent for backwards compatibility.
return rel == Equivalent || rel == Ancestor
case scopesv1.Mode_MODE_RELATIVES:
return rel != Orthogonal
default:
// unknown modes match nothing.
return false
}
}
// ValidateFilter checks that a filter is well-formed: that its mode is recognized and that its scope is
// consistent with its mode. Relational modes (EXACT/DESCENDANTS/ANCESTORS/RELATIVES) require a
// non-empty/valid scope, while the non-relational modes (UNSCOPED/ALL) and the unspecified mode
// require an empty scope. A nil or unspecified filter is considered valid.
func ValidateFilter(filter *scopesv1.Filter) error {
switch filter.GetMode() {
case scopesv1.Mode_MODE_UNSPECIFIED:
if filter.GetScope() != "" {
return trace.BadParameter("scope filter specifies a scope %q without a mode", filter.GetScope())
}
return nil
case scopesv1.Mode_MODE_EXACT,
scopesv1.Mode_MODE_DESCENDANTS,
scopesv1.Mode_MODE_ANCESTORS,
scopesv1.Mode_MODE_RELATIVES,
scopesv1.Mode_MODE_POLICIES_APPLICABLE_TO_SCOPE: //nolint:staticcheck // SA1019. Deprecated mode retained as equivalent for backwards compatibility.
if filter.GetScope() == "" {
return trace.BadParameter("scope filter mode %v requires a non-empty scope", filter.GetMode())
}
if err := WeakValidate(filter.GetScope()); err != nil {
return trace.Wrap(err, "invalid scope in scope filter")
}
return nil
case scopesv1.Mode_MODE_UNSCOPED, scopesv1.Mode_MODE_ALL:
if filter.GetScope() != "" {
return trace.BadParameter("scope filter mode %v requires an empty scope, got %q", filter.GetMode(), filter.GetScope())
}
return nil
default:
return trace.BadParameter("unknown scope filter mode %v", filter.GetMode())
}
}
/*
* Teleport
* Copyright (C) 2026 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 scopes
import (
"strings"
"github.com/gravitational/trace"
)
// QualifiedNameSeparator is the separator between scope and name in a
// scope-qualified name. This separator must never appear in scope segments.
const QualifiedNameSeparator = "::"
// QualifiedName pairs a scope with a resource name to uniquely identify a scoped
// resource. The canonical form of a scope-qualified name (SQN) is "<scope>::<name>",
// e.g. "/staging/west::myrole". SQNs take the place of bare names in configuration
// resource specs, CLI interfaces, and user-facing messages (e.g. errors) where it
// is necessary to fully specify the unique identifier of a scoped resource. Internally,
// teleport APIs should generally continue to use separate scope and name fields, as
// should structured logs/events.
//
// A QualifiedName may be used in APIs that need to refer to a resource that
// may be scoped or unscoped. In these cases the Scope field may be empty, and
// WeakValidate and StrongValidate will return an error.
type QualifiedName struct {
// Scope is the resource's scope path, e.g. "/staging/west".
Scope string
// Name is the resource's name within its scope, e.g. "myrole".
Name string
}
// String returns the string representation of the QualifiedName.
// If the Scope is empty, the Name is returned verbatim.
func (q QualifiedName) String() string {
if q.Scope == "" {
return q.Name
}
return q.Scope + QualifiedNameSeparator + q.Name
}
// Set sets a possible scope qualified name. Input that does not look like an
// SQN (see [MaybeSQN]) becomes a bare name with an empty scope. This implements
// the flag/kingping Value interface.
func (q *QualifiedName) Set(val string) error {
if !MaybeSQN(val) {
*q = QualifiedName{Name: val}
return nil
}
sqn, err := ParseQualifiedName(val)
if err != nil {
return err
}
if err := sqn.StrongValidate(); err != nil {
return err
}
*q = sqn
return nil
}
// StrongValidate validates this QualifiedName using strong validation rules. This method
// *must* be called on all QualifiedName values derived from user input and/or cluster-external
// sources. Use [QualifiedName.WeakValidate] when checking values from the control plane in
// logic that may run agent-side.
func (q QualifiedName) StrongValidate() error {
if err := StrongValidate(q.Scope); err != nil {
return trace.BadParameter("scope-qualified name %q has invalid scope: %v", q, err)
}
if err := StrongValidateResourceName(q.Name); err != nil {
return trace.BadParameter("scope-qualified name %q has invalid name: %v", q, err)
}
// as an extra precaution, also run all weak checks just to be certain we didn't accidentally
// construct a weak check that rejects something that would otherwise pass a strong check.
if err := q.WeakValidate(); err != nil {
return trace.Wrap(err)
}
return nil
}
// WeakValidate performs a weak form of validation on this QualifiedName. This is useful for
// ensuring that values received from trusted sources (e.g. the control plane) haven't been
// altered beyond our ability to reason effectively about them. Prefer [QualifiedName.StrongValidate]
// for values derived from external sources (e.g. user input).
func (q QualifiedName) WeakValidate() error {
if err := WeakValidate(q.Scope); err != nil {
return trace.BadParameter("scope-qualified name %q has invalid scope: %v", q, err)
}
if err := WeakValidateSegment(q.Name); err != nil {
return trace.BadParameter("scope-qualified name %q has invalid name: %v", q, err)
}
return nil
}
// MaybeSQN returns true if the given string *might* be a scope-qualified name. This function is intended to be used
// for testing fields that may contain a mix of scope-qualified and unscoped names. Generally, any string that trips
// this check should be considered to have been intended to be an SQN by the user, and treated as a typo if it fails
// to parse as one.
func MaybeSQN(s string) bool {
return strings.HasPrefix(s, separator) || strings.Contains(s, QualifiedNameSeparator)
}
// ParseQualifiedName parses a scope-qualified name string into its scope and name
// components by splitting on the first occurrence of "::". Returns an error if the
// separator is absent or either component is empty. This function does not validate
// the format of the scope or name components; use [QualifiedName.StrongValidate] or
// [QualifiedName.WeakValidate] for validation.
func ParseQualifiedName(sqn string) (QualifiedName, error) {
scope, name, ok := strings.Cut(sqn, QualifiedNameSeparator)
if !ok {
return QualifiedName{}, trace.BadParameter("scope-qualified name %q missing %q separator", sqn, QualifiedNameSeparator)
}
if scope == "" {
return QualifiedName{}, trace.BadParameter("scope-qualified name %q has empty scope component", sqn)
}
if name == "" {
return QualifiedName{}, trace.BadParameter("scope-qualified name %q has empty name component", sqn)
}
return QualifiedName{Scope: scope, Name: name}, nil
}
// StrongValidateQualifiedName validates a scope-qualified name string using strong validation
// rules. This function *must* be called on all scope-qualified name values received from
// user input and/or cluster-external sources. Use [WeakValidateQualifiedName] when
// checking values from the control plane in logic that may run agent-side.
//
// Prefer parsing with [ParseQualifiedName] and then calling [QualifiedName.StrongValidate]
// directly when the parsed value is needed, to avoid parsing twice.
func StrongValidateQualifiedName(sqn string) error {
qn, err := ParseQualifiedName(sqn)
if err != nil {
return trace.Wrap(err)
}
return qn.StrongValidate()
}
// WeakValidateQualifiedName performs a weak form of validation on a scope-qualified name string.
// This is useful for ensuring that values received from trusted sources (e.g. the control
// plane) haven't been altered beyond our ability to reason effectively about them. Prefer
// [StrongValidateQualifiedName] for values received from external sources (e.g. user input).
//
// Prefer parsing with [ParseQualifiedName] and then calling [QualifiedName.WeakValidate]
// directly when the parsed value is needed, to avoid parsing twice.
func WeakValidateQualifiedName(sqn string) error {
qn, err := ParseQualifiedName(sqn)
if err != nil {
return trace.Wrap(err)
}
return qn.WeakValidate()
}
/*
* Teleport
* Copyright (C) 2025 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 scopes
import (
"fmt"
"iter"
"regexp"
"strings"
"unicode"
"github.com/gravitational/trace"
)
// segmentRegexp is the regular expression used to validate scope segments. It allows
// lowercase alphanumeric characters, hyphens, underscores, and periods. It also requires
// that the segment starts and ends with an alphanumeric character.
var segmentRegexp = regexp.MustCompile(`^[a-z0-9][a-z0-9\-\_\.]*[a-z0-9]$`)
const (
// separator is the character used to separate segments in a scope and is the the value of the root scope.
separator = "/"
// exclusiveChildGlobSegment is the string used to indicate a wildcard match for child segments, used to construct
// the exclusive child glob suffix.
exclusiveChildGlobSegment = "**"
// exclusiveChildGlobSuffix is a special suffix used in roles to indicate that the role can be
// assigned to any *child* of a given scope, but not to the scope itself (e.g. an assignable scope of
// `/aa/**` allows assignment to `/aa/bb`, but not to `/aa`).
exclusiveChildGlobSuffix = separator + exclusiveChildGlobSegment
// maxScopeSize is the maximum size of a scope, including separators.
maxScopeSize = 64
// maxSegmentSize is the maximum size of a segment, excluding separators.
// The max size is 36 to take into account UUIDs.
maxSegmentSize = 36
// minSegmentSize is the minimum size of a segment, excluding separators.
minSegmentSize = 2
// breakingChars is a special set of characters that we explicitly consider to be invalid
// members even when performing looser/weaker validation. This it intended to be an additional
// guardrail against erroneous interpretation of future extensions to scope syntax by outdated
// agents in the event of improper cross-version compat logic.
breakingChars = "\\@(){}\"'%*?#$!+=|<>,;~`&[]/"
)
// Root is the root scope. Non-policy resources being grandfathered into scoping should use this
// as their default scope value, which will exclude the resource from administration by any policy
// other than non-scoped policies and root-scoped policies. Note that there is no sane default scoping
// for a policy. Policy resources should never use this or any other default scope value.
const Root = separator
// Unscoped is an empty string constant used in code to improve the readability of code that is intended
// to handle a mix of scoped and unscoped values.
const Unscoped = ""
// StrongValidate checks if the scope is valid according to all scope formatting rules. This function
// *must* be called on all scope values received from user input and/or cluster-external sources (e.g.
// an identity provider). Use of this function should be avoided when checking the validity of scopes
// from the control-plane in logic that may be run agent-side. Prefer [WeakValidate] in those cases, which
// is more forgiving of changes to scope formatting rules.
func StrongValidate(scope string) error {
if scope == "" {
return trace.BadParameter("scope is empty")
}
if !strings.HasPrefix(scope, separator) {
return trace.BadParameter("scope %q is missing required prefix %q", scope, separator)
}
if scope != separator && strings.HasSuffix(scope, separator) {
return trace.BadParameter("scope %q has dangling separator %q", scope, separator)
}
for segment := range DescendingSegments(scope) {
if err := StrongValidateSegment(segment); err != nil {
return trace.BadParameter("scope %q is invalid: %v", scope, err)
}
}
if len(scope) > maxScopeSize {
return trace.BadParameter("scope %q is too long (max characters %d)", scope, maxScopeSize)
}
// as an extra precaution, also run all weak checks just to be certain we didn't accidentally
// construct a weak check that rejects something that would otherwise pass a strong check. strong
// validation is not used in perf-critical paths, so there isn't any real downside to a little
// defensiveness here.
if err := WeakValidate(scope); err != nil {
return trace.BadParameter("scope would not pass weak validation: %v", err)
}
return nil
}
// WeakValidate performs a weak form of validation on a scope. This is useful primarily for ensuring
// that scopes received from trusted sources haven't been altered beyond our ability to reason effectively
// about them (e.g. due to significant version drift). Prefer using [StrongValidate] for scopes received from
// external sources (e.g. user input or identity provider).
func WeakValidate(scope string) error {
if scope == "" {
return trace.BadParameter("scope is empty")
}
for segment := range DescendingSegments(scope) {
if err := WeakValidateSegment(segment); err != nil {
return trace.BadParameter("scope %q is invalid: %v", scope, err)
}
}
return nil
}
// StrongValidateSegment checks if the scope segment is valid according to all scope formatting rules. This function
// *must* be called on all scope segment values received from user input and/or cluster-external sources (e.g.
// an identity provider). Use of this function should be avoided when checking the validity of segments
// from the control-plane in logic that may be run agent-side. Prefer [WeakValidateSegment] in those cases, which
// is more forgiving of changes to scope formatting rules.
func StrongValidateSegment(segment string) error {
if segment == "" {
return trace.BadParameter("segment is empty")
}
if len(segment) < minSegmentSize {
return trace.BadParameter("segment %q is too short (min characters %d)", segment, minSegmentSize)
}
if err := strongValidateFormat("segment", segment, segmentRegexp); err != nil {
return trace.Wrap(err)
}
if len(segment) > maxSegmentSize {
return trace.BadParameter("segment %q is too long (max characters %d)", segment, maxSegmentSize)
}
return nil
}
// strongValidateFormat applies the strong formatting checks to value, describing it as noun
// (e.g. "segment" or "name") in any error it returns:
// - no uppercase characters
// - the provided shape regexp
// - weak checks as a defensive backstop
func strongValidateFormat(noun, value string, shape *regexp.Regexp) error {
// check for uppercase characters separately. this would be caught by the regex, but its better
// UX to call out uppercase characters specifically since its a common mistake.
for _, r := range value {
if unicode.IsUpper(r) {
return trace.BadParameter("%s %q contains uppercase character(s)", noun, value)
}
}
if !shape.MatchString(value) {
return trace.BadParameter("%s %q is malformed", noun, value)
}
// as an extra precaution, also run all weak checks just to be certain we didn't accidentally
// construct a weak check that rejects something that would otherwise pass a strong check. strong
// validation is not used in perf-critical paths, so there isn't any real downside to a little
// defensiveness here.
if err := WeakValidateSegment(value); err != nil {
return trace.BadParameter("%s would not pass weak validation: %v", noun, err)
}
return nil
}
// WeakValidateSegment performs a weak form of validation on a scope segment. This is useful primarily for ensuring
// that segments received from trusted sources haven't been altered beyond our ability to reason effectively
// about them (e.g. due to significant version drift). Prefer using [StrongValidateSegment] for segments received from
// external sources (e.g. user input or identity provider).
func WeakValidateSegment(segment string) error {
if segment == "" {
return trace.BadParameter("segment is empty")
}
// check for spaces and control characters
for _, b := range []byte(segment) {
if !isNonSpacePrintableASCII(b) {
return trace.BadParameter("segment %q contains invalid character", segment)
}
}
// check for breaking characters
if strings.ContainsAny(segment, breakingChars) {
return trace.BadParameter("segment %q contains invalid character", segment)
}
return nil
}
// nameRegexp is the regular expression used to validate scoped resource names. It enforces
// the same character rules as segmentRegexp, but allows for single character names.
var nameRegexp = regexp.MustCompile(`^[a-z0-9]([a-z0-9\-\_\.]*[a-z0-9])?$`)
// StrongValidateResourceName checks if a scoped resource name is valid according to all resource name
// formatting rules. Scoped resource names follow the same character restrictions as scope segments, but
// are not subject to the maximum segment length limit. This function *must* be called on all scoped
// resource name values received from user input and/or cluster-external sources. Use of this function
// should be avoided when checking the validity of names from the control-plane in logic that may be run
// agent-side.
func StrongValidateResourceName(name string) error {
if name == "" {
return trace.BadParameter("name is empty")
}
return trace.Wrap(strongValidateFormat("name", name, nameRegexp))
}
// isNonSpacePrintableASCII checks if a byte is a non-space printable ASCII character (i.e. a byte in the range
// [33, 126] inclusive). This is used for weak validation of scope segments and globs.
func isNonSpacePrintableASCII(b byte) bool {
if b < 33 || b > 126 {
return false
}
return true
}
// StrongValidateGlob checks if the scope glob is valid according to all scope formatting rules. This function
// *must* be called on all scope glob values received from user input and/or cluster-external sources (e.g.
// an identity provider). Use of this function should be avoided when checking the validity of scope globs
// from the control-plane in logic that may be run agent-side. Prefer [WeakValidateGlob] in those cases, which
// is more forgiving of changes to scope glob formatting rules.
func StrongValidateGlob(scope string) error {
if scope == "" {
return trace.BadParameter("scope glob is empty")
}
if scope == exclusiveChildGlobSuffix {
// this is just a matcher for any child of root
return nil
}
if err := StrongValidate(strings.TrimSuffix(scope, exclusiveChildGlobSuffix)); err != nil {
return trace.BadParameter("scope glob %q is invalid: %v", scope, err)
}
// as an extra precaution, also run all weak checks just to be certain we didn't accidentally
// construct a weak check that rejects something that would otherwise pass a strong check. strong
// validation is not used in perf-critical paths, so there isn't any real downside to a little
// defensiveness here.
if err := WeakValidateGlob(scope); err != nil {
return trace.BadParameter("scope glob would not pass weak validation: %v", err)
}
return nil
}
// WeakValidateGlob is a weaker form of validation for scope globs. This is useful primarily for ensuring
// that scope globs received from trusted sources haven't been altered beyond our ability to reason effectively
// about them (e.g. due to significant version drift). Prefer using [StrongValidateGlob] for globs received from
// external sources (e.g. user input or identity provider).
func WeakValidateGlob(scope string) error {
if scope == "" {
return trace.BadParameter("scope glob is empty")
}
for segment := range DescendingSegments(scope) {
if segment == exclusiveChildGlobSegment {
continue
}
if err := WeakValidateSegment(segment); err != nil {
return trace.BadParameter("scope glob %q is invalid: %v", scope, err)
}
}
return nil
}
// DescendingSegments produces an iterator over the segments of a scope in descending order.
// e.g. `DescendingSegments("/a/b/c")` will result in an iterator that returns `a`, `b`, and
// `c` in that order. `DescendingSegments("/")` will return an empty iterator.
//
// Note that this function does not perform validation and is deliberately more relaxed about
// its inputs than our validation functions allow.
func DescendingSegments(scope string) iter.Seq[string] {
if scope == "" || scope == separator {
return func(yield func(string) bool) {}
}
return strings.SplitSeq(trimForSplit(scope), separator)
}
// Split splits the scope into its component segments and returns them as a slice.
// e.g. `Split("/a/b/c")` will return `[]string{"a", "b", "c"}`. `Split("/")` will return
// an empty slice.
//
// Note that this function does not perform validation and is deliberately more relaxed about
// its inputs than our validation functions allow.
func Split(scope string) []string {
if scope == "" || scope == separator {
return nil
}
return strings.Split(trimForSplit(scope), separator)
}
// Depth returns the depth of a scope (i.e., the number of segments) e.g. `Depth("/a/b/c")` returns 3,
// `Depth("/")` returns 0, etc. This is equivalent to `len(Split(scope))` but avoids allocation.
//
// Note that this function does not perform validation and is deliberately more relaxed about
// its inputs than our validation functions allow.
func Depth(scope string) int {
if scope == "" || scope == separator {
return 0
}
trimmed := trimForSplit(scope)
if trimmed == "" {
return 0
}
return strings.Count(trimmed, separator) + 1
}
func trimForSplit(scope string) string {
return strings.TrimSuffix(strings.TrimPrefix(scope, separator), separator)
}
// DescendingScopes produces an iterator over the canonical representations of the parents of
// a scope in descending order, ending with the canonical representation of the scope itself.
// e.g. `DescendingScopes("/a/b/c")` will result in an iterator that returns `/`, `/a`, `/a/b`,
// and `/a/b/c` in that order. `DescendingScopes("/")` will return a single-element iterator, and
// `DescendingScopes("")` will return an empty iterator.
//
// Note that this function does not perform validation and is deliberately more relaxed about
// its inputs than our validation functions allow.
func DescendingScopes(scope string) iter.Seq[string] {
return func(yield func(string) bool) {
if scope == "" {
return
}
// start with the root scope
if !yield(separator) {
return
}
var segments []string
// iterate over the segments in descending order
for segment := range DescendingSegments(scope) {
segments = append(segments, segment)
if !yield(Join(segments...)) {
return
}
}
}
}
// AscendingScopes produces an iterator over the canonical representations of a scope and its
// parents in ascending order, starting with the canonical representation of the scope itself
// and ending with the canonical representation of the root scope.
// e.g. `AscendingScopes("/a/b/c")` will result in an iterator that returns `/a/b/c`, `/a/b`, `/a`, and `/` in that order.
// // Note that this function does not perform validation and is deliberately more relaxed about
// its inputs than our validation functions allow.
func AscendingScopes(scope string) iter.Seq[string] {
return func(yield func(string) bool) {
if scope == "" {
return
}
segments := make([]string, 0, strings.Count(scope, separator))
// iterate over the segments in descending order, building the scope as we go
for segment := range DescendingSegments(scope) {
segments = append(segments, segment)
}
// now yield the scopes in ascending order
for i := len(segments); i >= 0; i-- {
if !yield(Join(segments[:i]...)) {
return
}
}
}
}
// Join joins the given segments into a single scope string. Note that this function
// does not perform validation and will produce invalid scopes if one or more segments
// are invalid.
func Join(segments ...string) string {
scope := separator + strings.Join(segments, separator)
// a tricky bit of how scope splitting/joining works is that an empty scope component
// as the last component is represented by `/aa//` not `/aa/`, because we chose not to assign trailing
// separators any meaning (i.e. `/aa/` does not imply the existence of an empty scope segment
// after `aa`). Trailing separators and empty scope segments are both considered invalid, but adding
// this behavior ensures that our splitting/joining logic cannot cannot change the meaning of a scope
// value when applied to one-another's outputs. If we don't add this logic, then
// DescendingSegments(Join([]string{"aa", ""}...)) will produce []string{"aa"}. We could just as easily amend
// `DescendingSegments` to emit an empty segment in the event of a trailing separator, but
// we judged that to be the higher risk behavior since that would effectively "move" a value to a new
// scope in the event of a trailing separator making it past validation. Our chosen behavior only "moves" an assignment
// in the event of a double-separator getting through, which is a much more visually obvious and intuitively incorrect
// value, and therefore less likely to be hit during processing of user input.
if len(segments) > 0 && segments[len(segments)-1] == "" {
scope += separator
}
return scope
}
// NormalizeForEquality attempts to normalize a scope s.t. equivalent scopes will have identical string representations. The only way
// scopes can currently be equivalent without being identical is by lacking a leading separator or containing a trailing separator.
// In general, we prefer to avoid modifying scope values as a verson incompatibility or bug that causes modification of a scope to change
// the result of scope comparison could result in a security issue. Normalizing on separators is the exception because separators
// are purely syntactic and have no semantic meaning beyond delimiting segments, and we effectively are already forced to normalize
// on separators (especially where scope caching is concerned, as scoped values must be organized by segment hierarchy for efficient
// lookups).
//
// NOTE: it generally holds that NormalizeForEquality(x) == NormalizeForEquality(x) if Compare(x, y) == Equivalent and vice versa, *except*
// in the case of two empty scopes. Empty scopes are invalid and considered orthogonal to one another, a proprerty which isn't really reasonable
// to preserve in this normalization function. Avoid using this function in contexts where scope values may be empty.
func NormalizeForEquality(scope string) string {
if scope == "" {
return ""
}
segments := make([]string, 0, strings.Count(scope, separator))
for segment := range DescendingSegments(scope) {
segments = append(segments, segment)
}
return Join(segments...)
}
// Relationship describes the relationship between two scopes, as determined by the [Compare] function. Note
// that direct use of this type in access-control logic is discouraged, as it is easier to accidentally
// misuse than the provided helpers (e.g. [ScopeOfOrigin]).
type Relationship int
const (
// Orthogonal indicates that the scopes are divergents/unrelated (e.g. '/foo' and '/bar').
Orthogonal Relationship = iota
// Equivalent indicates that the scopes are equal. Some non-equal scope strings are still
// considered equivalent by [Compare] (e.g. 'foo' and '/foo/'), though the canonical representations
// (e.g. '/foo') will also have string equality.
Equivalent
// Ancestor indicates that one scope is an ancestor of another (including multi-level ancestors,
// e.g. '/foo' is an ancestor of '/foo/bar' and '/foo/bar/bin', and '/' is an ancestor to all
// other scope values).
Ancestor
// Descendant indicates that one scope is a descendant of another (including multi-level descendants,
// e.g. '/foo/bar/bin' is a descendant of '/foo/bar' and '/foo', and all other scope values are
// descendants of '/').
Descendant
)
// String returns the human-readable representation of the relationship.
func (rel Relationship) String() string {
switch rel {
case Orthogonal:
return "Orthogonal"
case Equivalent:
return "Equivalent"
case Ancestor:
return "Ancestor"
case Descendant:
return "Descendant"
default:
return fmt.Sprintf("Unknown(%d)", rel)
}
}
// Compare compares the relationship between two scopes. The returned value is the
// relationship of the second scope to the first. I.e. if Compare(x, y) returns Ancestor,
// then y is an ancestor of x.
//
// Returns:
//
// Compare(/aa, /aa) => Equivalent
// Compare(/aa, /bb) => Orthogonal
// Compare(/aa, /aa/bb) => Descendant
// Compare(/aa/bb, /aa) => Ancestor
//
// Prefer using one of the provided helpers (e.g. [ScopeOfOrigin]) over using this function directly,
// as direct usage of this function can lead to ambiguity and accidental misuse.
//
// Note that this function does not perform validation, and may return unexpected results when
// called against invalid scope values.
func Compare(lhs, rhs string) Relationship {
if lhs == "" || rhs == "" {
// empty scopes are always orthogonal, including to one-another
return Orthogonal
}
lNext, lStop := iter.Pull(DescendingSegments(lhs))
defer lStop()
rNext, rStop := iter.Pull(DescendingSegments(rhs))
defer rStop()
for {
lVal, lOk := lNext()
rVal, rOk := rNext()
switch {
case lOk && rOk:
// both scopes have segments left to compare
if lVal == rVal {
// scopes are still equivalent at this level, continue processing
continue
}
// scopes have diverged
return Orthogonal
case lOk && !rOk:
// the right hand side scope is an ancestor of the left hand side scope
return Ancestor
case !lOk && rOk:
// the left hand side scope is an ancestor of the right hand side scope
return Descendant
case !lOk && !rOk:
// scopes are equivalent
return Equivalent
}
}
}
// Sort is a helper function for sorting scopes. Scope sort order differs from lexographic sort of
// a scope's string representation in some cases. For example, the lexicographically sorted sequence
// of scopes ['/aa', '/aa-bb', '/aa/bb'] is different from the scope-sorted sequence
// ['/aa', '/aa/bb', '/aa-bb']. This function conforms to the standard go comparison contract, returning
// a negative integer if lhs < rhs, zero if lhs == rhs, and a positive integer if lhs > rhs.
func Sort(lhs, rhs string) int {
if lhs == rhs {
return 0
}
lNext, lStop := iter.Pull(DescendingSegments(lhs))
defer lStop()
rNext, rStop := iter.Pull(DescendingSegments(rhs))
defer rStop()
for {
lVal, lOk := lNext()
rVal, rOk := rNext()
switch {
case lOk && rOk:
// both scopes have segments left to compare
if c := strings.Compare(lVal, rVal); c == 0 {
// scopes are still equivalent at this level, continue processing
continue
} else {
// scopes have diverged
return c
}
case lOk && !rOk:
// the right hand side scope is an ancestor of the left hand side scope
return 1
case !lOk && rOk:
// the left hand side scope is an ancestor of the right hand side scope
return -1
case !lOk && !rOk:
// scopes are equivalent
return 0
}
}
}
// ScopeOfOrigin is a helper for constructing unambiguous checks in access control logic. Prefer helpers like
// this over using the Compare function directly, as it improves readability and reduces the risk of misuse. Ex:
//
// if scopes.ScopeOfOrigin(roleScope).IsAssignableToScopeOfEffect(assignmentScope) { ... }
//
// Note that this helper does not perform validation, and may produce unexpected results when used against
// invalid scope values.
type ScopeOfOrigin string
// IsAssignableToScopeOfEffect checks if the Scope of Origin is compatible with the specified Scope of Effect. More
// specifically, this method returns true if a policy originating from this Scope of Origin is able to be assigned
// to the provided Scope of Effect. Policies may only have Scopes of Effect that are equivalent or descendant. A policy
// originating from a child scope being effectual in a parent scope would violate scope isolation principles.
func (s ScopeOfOrigin) IsAssignableToScopeOfEffect(scope string) bool {
rel := Compare(string(s), scope)
return rel == Equivalent || rel == Descendant
}
// ScopeOfEffect is a helper for constructing unambiguous checks in access control logic. Prefer helpers like
// this over using the Compare function directly, as it improves readability and reduces the risk of misuse. Ex:
//
// if scopes.ScopeOfEffect(assignment.ScopeOfEffect).IsAssignableFromScopeOfOrigin(assignment.ScopeOfOrigin) { ... }
//
// if scopes.ScopeOfEffect(assignment.ScopeOfEffect).AppliesToResourceScope(node.Scope) { ... }
//
// Note that this helper does not perform validation, and may produce unexpected results when used against
// invalid scope values.
type ScopeOfEffect string
// IsAssignableFromScopeOfOrigin checks if the Scope of Effect is compatible with the specified Scope of Origin. See
// [ScopeOfOrigin.IsAssignableToScopeOfEffect] for discussion of the importance of this check.
func (s ScopeOfEffect) IsAssignableFromScopeOfOrigin(scope string) bool {
rel := Compare(string(s), scope)
return rel == Equivalent || rel == Ancestor
}
// AppliesToResourceScope checks if this scope of effect applies to the specified resource scope. See [ResourceScope.IsSubjectToScopeOfEffect]
// for discussion of the importance of this check.
func (s ScopeOfEffect) AppliesToResourceScope(scope string) bool {
rel := Compare(string(s), scope)
return rel == Equivalent || rel == Descendant
}
// ResourceScope is a helper for constructing unambiguous checks in access control logic. Prefer helpers like
// this over using the Compare function directly, as it improves readability and reduces the risk of misuse. Ex:
//
// if scopes.ResourceScope(nodeScope).IsSubjectToScopeOfEffect(roleAssignmentScope) { ... }
//
// Note that this helper does not perform validation, and may produce unexpected results when used against
// invalid scope values.
type ResourceScope string
// IsSubjectToScopeOfEffect checks if this resource scope is subject to the specified policy Scope of Effect. Scoped
// policies always apply only to a given scope and its descendants. This method tells us if the resource scope in
// question falls under the purview of the specified Scope of Effect. Policies whose Scope of Effect does not apply
// to a resource must have no effect on that resource in order to preserve scope isolation.
func (s ResourceScope) IsSubjectToScopeOfEffect(scope string) bool {
rel := Compare(string(s), scope)
return rel == Equivalent || rel == Ancestor
}
// PolicyResourceScope is a helper for constructing unambiguous checks in access control logic. Prefer helpers like
// this over using the Compare function directly, as it improves readability and reduces the risk of misuse. Ex:
//
// if scopes.PolicyResourceScope(assignment.Scope).CanDependOnStateFromPolicyResourceAtScope(role.Scope) { ... }
//
// Note that this helper does not perform validation, and may produce unexpected results when used against
// invalid scope values.
type PolicyResourceScope string
// CanDependOnStateFromPolicyResourceAtScope checks if this policy resource scope can depend on state from a policy
// resource at the specified scope. Policy state must always flow from ancestor/equivalent scopes to descendant/equivalent
// scopes (e.g. a role assignment can only reference roles in its scope or its ancestor scopes).
func (s PolicyResourceScope) CanDependOnStateFromPolicyResourceAtScope(scope string) bool {
rel := Compare(string(s), scope)
return rel == Equivalent || rel == Ancestor
}
// ScopeOfEffectGlob is a helper for constructing unambiguous checks in access control logic. Prefer helpers like
// this over using the Glob type directly, as it improves readability and reduces the risk of misuse. Ex:
//
// if scopes.ScopeOfEffectGlob(role.Spec.AssignableScopes[0]).MatchesScopeOfEffect(assignment.ScopeOfEffect) { ... }
//
// Note that this helper does not perform validation, and may produce unexpected results when used against
// invalid glob values.
type ScopeOfEffectGlob string
// MatchesScopeOfEffectLiteral checks if the Scope of Effect glob matches the specified Scope of Effect literal. If the assignability
// of a policy is constrained by a ScopeOfEffectGlob, this function can be used to determine if a given Scope of Effect will
// be compatible with that constraint.
func (s ScopeOfEffectGlob) MatchesScopeOfEffectLiteral(scope string) bool {
return Glob(s).MatchesScopeLiteral(scope)
}
// IsAlwaysAssignableFromScopeOfOrigin checks if the Scope of Effect glob will exclusively match Scope of Effect literals
// that are considered assignable from the specified Scope of Origin. Policy resources such as scoped roles sometimes constrain
// the scopes they are assignable to using globs. This function ensures that a glob being used in this manner will never match a
// Scope of Effect that would be invalid given the Scope of Origin of the policy.
func (s ScopeOfEffectGlob) IsAlwaysAssignableFromScopeOfOrigin(scope string) bool {
return Glob(s).OnlyMatchesSubjectsOf(scope)
}
// Glob is a helper for matching scope globs against scopes. This is currently used to support exactly
// one piece of special syntax, the use of `/component/**` to indicate that a role can be assigned to any child of
// the specified scope, but not to the scope itself. Ex:
//
// if scopes.Glob(assignableScope).Matches(assignmentScope) { ... }
//
// Note that this helper does not perform validation, and may produce unexpected results when used against
// invalid scope or glob values.
type Glob string
// MatchesScopeLiteral checks if the given scope literal matches this scope glob.
func (s Glob) MatchesScopeLiteral(scope string) bool {
return matchGlob(string(s), scope)
}
// OnlyMatchesSubjectsOf checks if this scope glob only matches scopes that are subjects of the specified scope literal. In this
// case "subjects of" means "descendants of or equivalent to".
func (s Glob) OnlyMatchesSubjectsOf(scope string) bool {
return globOnlyMatchesSubjectsOfScope(string(s), scope)
}
// matchGlob implements glob matching. note that this function technically supports some limited
// matching behaviors that we don't actually currently allow the use of. e.g. `/foo/**/bar` would match
// `/foo/baz/bar` (but not `/foo/baz/bin/bar` or `/foo/bar/`), but we only permit the use of the double-wildcard
// syntax in the trailing segment for simplicity's sake.
func matchGlob(glob string, scope string) bool {
if glob == "" || scope == "" {
return false
}
gNext, gStop := iter.Pull(DescendingSegments(glob))
defer gStop()
sNext, sStop := iter.Pull(DescendingSegments(scope))
defer sStop()
for {
gVal, gOk := gNext()
sVal, sOk := sNext()
switch {
case gOk && sOk:
// both values have segments left to compare
if gVal == sVal {
// segments are equivalent, descend into the next segment
continue
}
if gVal == exclusiveChildGlobSegment {
// double-wildcard matches any segment, continue
continue
}
// scopes have diverged
return false
case gOk && !sOk:
// the scope is an ancestor of the glob
return false
case !gOk && sOk:
// the glob is an ancestor of the scope
return true
case !gOk && !sOk:
// literal match
return true
}
}
}
// globOnlyMatchesSubjectsOfScope checks if the given glob only matches scopes that are subjects of the specified scope
// literal. In this case "subjects of" means "descendants of or equivalent to".
func globOnlyMatchesSubjectsOfScope(glob string, scope string) bool {
if glob == "" || scope == "" {
return false
}
gNext, gStop := iter.Pull(DescendingSegments(glob))
defer gStop()
sNext, sStop := iter.Pull(DescendingSegments(scope))
defer sStop()
for {
gVal, gOk := gNext()
sVal, sOk := sNext()
switch {
case gOk && sOk:
// both values have segments left to compare
if gVal == sVal {
// segments are equivalent, descend into the next segment
continue
}
if gVal == exclusiveChildGlobSegment {
// we've hit the first wildcard segment and are still descending
// through the segments of the scope. this means that the glob may
// match a scope orthogonal to the resource scope and therefore does
// not conform to subjugation rules.
return false
}
// scopes have diverged
return false
case gOk && !sOk:
// the scope is an ancestor of the glob, anything the glob matches
// will be subject to the scope.
return true
case !gOk && sOk:
// the glob is an ancestor of the scope, and is therefore trivially not
// subject to the scope.
return false
case !gOk && !sOk:
// literal match, the glob is subject by equivalence to the scope.
return true
}
}
}
// EnforcementPoint combines a Scope of Origin and Scope of Effect to define specific point or target that
// one or more policies may be defined for. Since each policy has some scope it originates *from* and some
// scope it applies *to*, any given scoped policy (e.g. a scoped role) should exist "at" a specific enforcement
// point during access-control evaluation. Policy evaluation order and precedence is determined primarily by the
// enforcement point at which the policy exists, with policies at more ancestral origin scopes taking precedence
// over those at more descendant origin scopes, and within a given origin scope, policies at more specific effect
// scopes taking precedence over those at more general effect scopes.
//
// This type is primarily intended to be used in conjunction with [EnforcementPointsForResourceScope] to
// help determine ordering of policy evaluation during access checks.
type EnforcementPoint struct {
// ScopeOfOrigin is the scope from which a policy originates. This represents the authority/provenance
// of the policy. Policies with more ancestral Scopes of Origin take precedence.
ScopeOfOrigin string
// ScopeOfEffect is the scope at which a policy's effects apply. Within a given Scope of Origin,
// policies with more descendant/specific Scopes of Effect take precedence.
ScopeOfEffect string
}
// EnforcementPointsForResourceScope yields all (ScopeOfOrigin, ScopeOfEffect) pairs that should be evaluated when
// checking access to a resource at the given scope. The pairs are yielded in evaluation order:
// - First by Scope of Origin (root to resourceScope, preserving scope hierarchy)
// - Then by Scope of Effect (resourceScope to ScopeOfOrigin, most specific first)
//
// This iterator is intended to be the core building block for higher-level policy evaluation logic. Policy
// evaluation logic should use this iterator to determine the primary order in which policies should be evaluated
// (though it should be noted that multiple policies may exist at any given enforcement point).
//
// Ex: EnforcementPointsForResourceScope("/staging/west") yields:
// - (/, /staging/west) - root origin, specific effect
// - (/, /staging) - root origin, less specific effect
// - (/, /) - root origin, root effect
// - (/staging, /staging/west) - staging origin, specific effect
// - (/staging, /staging) - staging origin, staging effect
// - (/staging/west, /staging/west) - west origin, west effect
func EnforcementPointsForResourceScope(resourceScope string) iter.Seq[EnforcementPoint] {
return func(yield func(EnforcementPoint) bool) {
// iterate through all Scopes of Origin from root to resourceScope
for scopeOfOrigin := range DescendingScopes(resourceScope) {
// For each Scope of Origin, iterate through Scopes of Effect from resourceScope
// UP to (and including) the current Scope of Origin. This ensures more specific
// policies are evaluated before more general ones within each origin level.
// We use AscendingScopes to go from most specific (resourceScope) to least specific (scopeOfOrigin).
for scopeOfEffect := range AscendingScopes(resourceScope) {
// Only yield if the Scope of Origin permits assignment to this Scope of Effect.
// A scope of effect cannot be more ancestral than its origin (that would violate
// the rule that policies cannot reach up to affect parent scopes).
if !ScopeOfOrigin(scopeOfOrigin).IsAssignableToScopeOfEffect(scopeOfEffect) {
// we've reached scopes of effect that are too ancestral for this origin
break
}
if !yield(EnforcementPoint{
ScopeOfOrigin: scopeOfOrigin,
ScopeOfEffect: scopeOfEffect,
}) {
return
}
}
}
}
}
/*
* 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"
"fmt"
"iter"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend"
)
// LockGetter is a service that gets locks.
type LockGetter interface {
// GetLock gets a lock by name.
GetLock(ctx context.Context, name string) (types.Lock, error)
// GetLocks gets all/in-force locks that match at least one of the targets when specified.
GetLocks(ctx context.Context, inForceOnly bool, targets ...types.LockTarget) ([]types.Lock, error)
ListLocks(ctx context.Context, limit int, startKey string, filter *types.LockFilter) ([]types.Lock, string, error)
RangeLocks(ctx context.Context, start, end string, filter *types.LockFilter) iter.Seq2[types.Lock, error]
}
// Access service manages roles and permissions.
type Access interface {
// GetRoles returns a list of roles.
GetRoles(ctx context.Context) ([]types.Role, error)
// ListRoles is a paginated role getter.
ListRoles(ctx context.Context, req *proto.ListRolesRequest) (*proto.ListRolesResponse, error)
// CreateRole creates a role.
CreateRole(ctx context.Context, role types.Role) (types.Role, error)
// UpdateRole updates an existing role.
UpdateRole(ctx context.Context, role types.Role) (types.Role, error)
// UpsertRole creates or updates role.
UpsertRole(ctx context.Context, role types.Role) (types.Role, error)
// GetRole returns role by name.
GetRole(ctx context.Context, name string) (types.Role, error)
// DeleteRole deletes role by name.
DeleteRole(ctx context.Context, name string) error
LockGetter
// UpsertLock upserts a lock.
UpsertLock(context.Context, types.Lock) error
// DeleteLock deletes a lock.
DeleteLock(context.Context, string) error
// ReplaceRemoteLocks replaces the set of locks associated with a remote cluster.
ReplaceRemoteLocks(ctx context.Context, clusterName string, locks []types.Lock) error
}
// AccessInternal extends the Access interface with auth-specific internal methods.
type AccessInternal interface {
Access
// AppendPutRoleActions adds conditional actions to an atomic write to create
// or update a role.
AppendPutRoleActions(
actions []backend.ConditionalAction,
role types.Role,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteRoleActions adds conditional actions to an atomic write to
// delete a role.
AppendDeleteRoleActions(
actions []backend.ConditionalAction,
name string,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
var dynamicLabelsErrorMessage = fmt.Sprintf("labels with %q prefix are not allowed in deny rules", types.TeleportDynamicLabelPrefix)
// CheckDynamicLabelsInDenyRules checks if any deny rules in the given role use
// labels prefixed with "dynamic/".
func CheckDynamicLabelsInDenyRules(r types.Role) error {
for _, kind := range types.LabelMatcherKinds {
labelMatchers, err := r.GetLabelMatchers(types.Deny, kind)
if err != nil {
return trace.Wrap(err)
}
for label := range labelMatchers.Labels {
if strings.HasPrefix(label, types.TeleportDynamicLabelPrefix) {
return trace.BadParameter("%s", dynamicLabelsErrorMessage)
}
}
const expressionMatch = `"` + types.TeleportDynamicLabelPrefix
if strings.Contains(labelMatchers.Expression, expressionMatch) {
return trace.BadParameter("%s", dynamicLabelsErrorMessage)
}
}
for _, where := range []string{
r.GetAccessReviewConditions(types.Deny).Where,
r.GetImpersonateConditions(types.Deny).Where,
} {
if strings.Contains(where, types.TeleportDynamicLabelPrefix) {
return trace.BadParameter("%s", dynamicLabelsErrorMessage)
}
}
return nil
}
/*
* 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 (
"cmp"
"context"
"fmt"
"log/slog"
"maps"
"net"
"slices"
"strings"
"time"
"github.com/gravitational/trace"
"k8s.io/apimachinery/pkg/runtime/schema"
"github.com/gravitational/teleport/api/constants"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/wrappers"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/set"
)
// AccessChecker interface checks access to resources based on roles, traits,
// and allowed resources
type AccessChecker interface {
// HasRole checks if the checker includes the role
HasRole(role string) bool
// RoleNames returns a list of role names
RoleNames() []string
// Traits returns the set of user traits
Traits() wrappers.Traits
// Roles returns the list underlying roles this AccessChecker is based on.
Roles() []types.Role
// CheckAccess checks access to the specified resource.
CheckAccess(r AccessCheckable, state AccessState, matchers ...RoleMatcher) error
// CheckConditionalAccess checks conditional access to the specified resource. If access is granted, it returns
// preconditions that must be satisfied. If access is denied, it returns an error. An empty list of preconditions
// and a nil error indicates that no additional preconditions are required for access.
CheckConditionalAccess(r AccessCheckable, state AccessState, matchers ...RoleMatcher) ([]*decisionpb.Precondition, error)
// CheckAccessToRemoteCluster checks access to remote cluster
CheckAccessToRemoteCluster(cluster types.RemoteCluster) error
// CheckAccessToRule checks access to a rule within a namespace.
CheckAccessToRule(context RuleContext, namespace string, rule string, verb string) error
// GuessIfAccessIsPossible guesses if access is possible for an entire category
// of resources.
// It responds the question: "is it possible that there is a resource of this
// kind that the current user can access?".
// GuessIfAccessIsPossible is used, mainly, for UI decisions ("should the tab
// for resource X appear"?). Most callers should use CheckAccessToRule instead.
GuessIfAccessIsPossible(ctx RuleContext, namespace string, resource string, verb string) error
// CheckLoginDuration checks if role set can login up to given duration and
// returns a combined list of allowed logins.
CheckLoginDuration(ttl time.Duration) ([]string, error)
// CheckKubeGroupsAndUsers check if role can login into kubernetes
// and returns two lists of combined allowed groups and users
CheckKubeGroupsAndUsers(ttl time.Duration, overrideTTL bool, matchers ...RoleMatcher) (groups []string, users []string, err error)
// CheckAWSRoleARNs returns a list of AWS role ARNs role is allowed to assume.
CheckAWSRoleARNs(ttl time.Duration, overrideTTL bool) ([]string, error)
// CheckAzureIdentities returns a list of Azure identities the user is allowed to assume.
CheckAzureIdentities(ttl time.Duration, overrideTTL bool) ([]string, error)
// CheckGCPServiceAccounts returns a list of GCP service accounts the user is allowed to assume.
CheckGCPServiceAccounts(ttl time.Duration, overrideTTL bool) ([]string, error)
// CheckAccessToSAMLIdP checks access to SAML IdP service provider resource.
// It checks for both the legacy RBAC (role v7 and below) that checks for IDP
// role option and MFA, as well as non-legacy RBAC (role v8 and above) that checks
// for labels, MFA and Device Trust.
CheckAccessToSAMLIdP(r AccessCheckable, authPref readonly.AuthPreference, state AccessState, matchers ...RoleMatcher) error
// AdjustSessionTTL will reduce the requested ttl to lowest max allowed TTL
// for this role set, otherwise it returns ttl unchanged
AdjustSessionTTL(ttl time.Duration) time.Duration
// AdjustClientIdleTimeout adjusts requested idle timeout
// to the lowest max allowed timeout, the most restrictive
// option will be picked
AdjustClientIdleTimeout(ttl time.Duration) time.Duration
// AdjustDisconnectExpiredCert adjusts the value based on the role set
// the most restrictive option will be picked
AdjustDisconnectExpiredCert(disconnect bool) bool
// CheckAgentForward checks if the role can request agent forward for this
// user.
CheckAgentForward(login string) error
// CanForwardAgents returns true if this role set offers capability to forward
// agents.
CanForwardAgents() bool
// CanPortForward returns true if this RoleSet can forward ports.
CanPortForward() bool
// SSHPortForwardMode returns the SSHPortForwardMode that the RoleSet allows.
SSHPortForwardMode() decisionpb.SSHPortForwardMode
// DesktopClipboard returns true if the role set has enabled shared
// clipboard for desktop sessions. Clipboard sharing is disabled if
// one or more of the roles in the set has disabled it.
DesktopClipboard() bool
// RecordDesktopSession returns true if a role in the role set has enabled
// desktop session recoring.
RecordDesktopSession() bool
// DesktopDirectorySharing returns true if the role set has directory sharing
// enabled. This setting is enabled if one or more of the roles in the set has
// enabled it.
DesktopDirectorySharing() bool
// MaybeCanReviewRequests attempts to guess if this RoleSet belongs
// to a user who should be submitting access reviews. Because not all rolesets
// are derived from statically assigned roles, this may return false positives.
MaybeCanReviewRequests() bool
// PermitX11Forwarding returns true if this RoleSet allows X11 Forwarding.
PermitX11Forwarding() bool
// CanCopyFiles returns true if the role set has enabled remote file
// operations via SCP or SFTP. Remote file operations are disabled if
// one or more of the roles in the set has disabled it.
CanCopyFiles() bool
// CertificateFormat returns the most permissive certificate format in a
// RoleSet.
CertificateFormat() string
// EnhancedRecordingSet returns a set of events that will be recorded
// for enhanced session recording.
EnhancedRecordingSet() map[string]bool
// CheckDatabaseNamesAndUsers returns database names and users this role
// is allowed to use.
CheckDatabaseNamesAndUsers(ttl time.Duration, overrideTTL bool) (names []string, users []string, err error)
// DatabaseAutoUserMode returns whether a user should be auto-created in
// the database.
DatabaseAutoUserMode(types.Database) (types.CreateDatabaseUserMode, error)
// CheckDatabaseRoles returns a list of database roles to assign, when
// auto-user provisioning is enabled. If no user-requested roles, all
// allowed roles are returned.
CheckDatabaseRoles(database types.Database, userRequestedRoles []string) (roles []string, err error)
// GetDatabasePermissions returns a set of database permissions applicable for the user.
GetDatabasePermissions(database types.Database) (allow types.DatabasePermissions, deny types.DatabasePermissions, err error)
// CheckImpersonate checks whether current user is allowed to impersonate
// users and roles
CheckImpersonate(currentUser, impersonateUser types.User, impersonateRoles []types.Role) error
// CheckImpersonateRoles checks whether the current user is allowed to
// perform roles-only impersonation.
CheckImpersonateRoles(currentUser types.User, impersonateRoles []types.Role) error
// CanImpersonateSomeone returns true if this checker has any impersonation rules
CanImpersonateSomeone() bool
// LockingMode returns the locking mode to apply with this checker.
LockingMode(defaultMode constants.LockingMode) constants.LockingMode
// ExtractConditionForIdentifier returns a restrictive filter expression
// for list queries based on the rules' `where` conditions.
ExtractConditionForIdentifier(ctx RuleContext, namespace, resource, verb, identifier string) (*types.WhereExpr, error)
// CertificateExtensions returns the list of extensions for each role in the RoleSet
CertificateExtensions() []*types.CertExtension
// GetAllowedSearchAsRoles returns all of the allowed SearchAsRoles.
GetAllowedSearchAsRoles(allowFilters ...SearchAsRolesOption) []string
// GetAllowedSearchAsRolesForKubeResourceKind returns all of the allowed SearchAsRoles
// that allowed requesting to the requested Kubernetes resource kind.
GetAllowedSearchAsRolesForKubeResourceKind(requestedKubeResourceKind string) []string
// GetAllowedPreviewAsRoles returns all of the allowed PreviewAsRoles.
GetAllowedPreviewAsRoles() []string
// CheckSubmitForUser checks whether the current user is allowed to
// submit reviews for other users, to be used by plugins.
CheckSubmitForUser(currentUser, submitForUser types.User) error
// MaxConnections returns the maximum number of concurrent ssh connections
// allowed. If MaxConnections is zero then no maximum was defined and the
// number of concurrent connections is unconstrained.
MaxConnections() int64
// MaxSessions returns the maximum number of concurrent ssh sessions per
// connection. If MaxSessions is zero then no maximum was defined and the
// number of sessions is unconstrained.
MaxSessions() int64
// SessionPolicySets returns the list of SessionPolicySets for all roles.
SessionPolicySets() []*types.SessionTrackerPolicySet
// GetAllLogins returns all valid unix logins for the AccessChecker.
GetAllLogins() []string
// GetAllowedResourceAccessIDs returns the list of allowed resources the identity for
// the AccessChecker is allowed to access. An empty or nil list indicates that
// there are no resource-specific restrictions.
GetAllowedResourceAccessIDs() []types.ResourceAccessID
// SessionRecordingMode returns the recording mode for a specific service.
SessionRecordingMode(service constants.SessionRecordingService) constants.SessionRecordingMode
// HostUsers returns host user information matching a server or nil if
// a role disallows host user creation
HostUsers(types.Server) (*HostUsersDecision, error)
// HostSudoers returns host sudoers entries matching a server
HostSudoers(types.Server) ([]string, error)
// DesktopGroups returns the desktop groups a user is allowed to create or an access denied error if a role disallows desktop user creation
DesktopGroups(types.WindowsDesktop) ([]string, error)
// PinSourceIP forces the same client IP for certificate generation and SSH usage
PinSourceIP() bool
// GetAccessState returns the AccessState for the user given their roles, the
// cluster auth preference, and whether MFA and the user's device were
// verified.
GetAccessState(authPref readonly.AuthPreference) AccessState
// PrivateKeyPolicy returns the enforced private key policy for this role set,
// or the provided defaultPolicy - whichever is stricter.
PrivateKeyPolicy(defaultPolicy keys.PrivateKeyPolicy) (keys.PrivateKeyPolicy, error)
// GetKubeResources returns the allowed and denied Kubernetes Resources configured
// for a user.
GetKubeResources(cluster types.KubeCluster) (allowed, denied []types.KubernetesResource)
// EnumerateEntities works on a given role set to return a minimal description
// of allowed set of entities (db_users, db_names, etc). It is biased towards
// *allowed* entities; It is meant to describe what the user can do, rather than
// cannot do. For that reason if the user isn't allowed to pick *any* entities,
// the output will be empty.
//
// In cases where * is listed in set of allowed entities, it may be hard for
// users to figure out the expected entity to use. For this reason the parameter
// extraEntities provides an extra set of entities to be checked against
// RoleSet. This extra set of entities may be sourced e.g. from user connection
// history.
EnumerateEntities(resource AccessCheckable, listFn roleEntitiesListFn, newMatcher roleMatcherFactoryFn, extraEntities ...string) EnumerationResult
// EnumerateDatabaseUsers specializes EnumerateEntities to enumerate db_users.
EnumerateDatabaseUsers(database types.Database, extraUsers ...string) (EnumerationResult, error)
// EnumerateDatabaseNames specializes EnumerateEntities to enumerate db_names.
EnumerateDatabaseNames(database types.Database, extraNames ...string) EnumerationResult
// EnumerateMCPTools specializes EnumerateEntities to enumerate mcp.tools.
// mcp.tools support regexes and blobs so those expressions are returned.
EnumerateMCPTools(app types.Application) EnumerationResult
// GetAllowedLoginsForResource returns all of the allowed logins for the passed resource.
//
// Supports the following resource types:
//
// - types.Server with GetKind() == types.KindNode
// - types.KindWindowsDesktop
// - types.KindApp with IsAWSConsole() == true
GetAllowedLoginsForResource(resource AccessCheckable) ([]string, error)
// CheckSPIFFESVID checks if the role set has access to generating the
// requested SPIFFE ID. Returns an error if the role set does not have the
// ability to generate the requested SVID.
CheckSPIFFESVID(spiffeIDPath string, dnsSANs []string, ipSANs []net.IP) error
// AccessInfo returns the AccessInfo that this access checker is based on.
AccessInfo() *AccessInfo
// DelegationSessionID returns the ID of the current Delegation Session.
DelegationSessionID() string
// MaxKubernetesConnections returns the maximum number of concurrent
// Kubernetes connections allowed. If MaxKubernetesConnections is zero then
// no maximum was defined and the number of concurrent connections is
// unconstrained.
MaxKubernetesConnections() int64
}
// AccessInfo hold information about an identity necessary to check whether that
// identity has access to cluster resources. This info can come from a user or
// host SSH certificate, TLS certificate, or user information stored in the
// backend.
type AccessInfo struct {
// ScopePin is an optional pin that ties an identity to a specific scope and set of scoped roles. When
// set, the Roles field must not be set.
ScopePin *scopesv1.Pin
// Roles is the list of cluster local roles for the identity.
Roles []string
// Traits is the set of traits for the identity.
Traits wrappers.Traits
// AllowedResourceAccessIDs is the list of resource IDs the identity is allowed to
// access. A nil or empty list indicates that no resource-specific
// access restrictions should be applied. Used for search-based access
// requests.
AllowedResourceAccessIDs []types.ResourceAccessID
// DelegationSessionID is the ID of the Delegation Session this identity was
// created for, if any.
DelegationSessionID string
// Username is the Teleport username.
Username string
}
// accessChecker implements the AccessChecker interface.
type accessChecker struct {
info *AccessInfo
localCluster string
// RoleSet is embedded to use the existing implementation for most
// AccessChecker methods. Methods which require AllowedResourceAccessIDs (relevant
// to search-based access requests) will be implemented by
// accessChecker.
RoleSet
}
// NewAccessChecker returns a new AccessChecker which can be used to check
// access to resources.
// Args:
// - `info *AccessInfo` should hold the roles, traits, and allowed resource IDs
// for the identity.
// - `localCluster string` should be the name of the local cluster in which
// access will be checked. You cannot check for access to resources in remote
// clusters.
// - `access RoleGetter` should be a RoleGetter which will be used to fetch the
// full RoleSet
func NewAccessChecker(info *AccessInfo, localCluster string, access RoleGetter) (AccessChecker, error) {
if info.ScopePin != nil {
return nil, trace.Errorf("cannot create standard access checker: %w", ErrScopedIdentity)
}
roleSet, err := FetchRolesWithContext(info.Roles, access, RoleTemplateContext{
Username: info.Username,
Traits: info.Traits,
})
if err != nil {
return nil, trace.Wrap(err)
}
return newAccessChecker(info, localCluster, roleSet), nil
}
// NewAccessCheckerForUserSession is an alternative to NewAccessChecker that includes a UserSessionRoleNotFoundErrorMsg if
// a role from the user's session is not found during the access check. This allows the Web UI to distinguish between
// a user session role lookup error (which should prompt the user to re-login) vs. other role lookup
// failures.
func NewAccessCheckerForUserSession(info *AccessInfo, localCluster string, access RoleGetter) (AccessChecker, error) {
roleSet, err := FetchRolesWithContext(info.Roles, access, RoleTemplateContext{
Username: info.Username,
Traits: info.Traits,
})
if err != nil {
if trace.IsNotFound(err) {
// Add the UserSessionRoleNotFoundErrorMsg message to indicate this role not found error was encountered fetching
// the user's session roles. This can only happen if the user's session certificate contains a role that no longer exists.
return nil, trace.Wrap(err, UserSessionRoleNotFoundErrorMsg)
}
return nil, trace.Wrap(err)
}
return &accessChecker{
info: info,
localCluster: localCluster,
RoleSet: roleSet,
}, nil
}
// NewAccessCheckerWithRoleSet is similar to NewAccessChecker, but accepts the
// full RoleSet rather than a RoleGetter.
func NewAccessCheckerWithRoleSet(info *AccessInfo, localCluster string, roleSet RoleSet) AccessChecker {
return newAccessChecker(info, localCluster, roleSet)
}
func newAccessChecker(info *AccessInfo, localCluster string, roleSet RoleSet) *accessChecker {
return &accessChecker{
info: info,
localCluster: localCluster,
RoleSet: roleSet,
}
}
// CurrentUserRoleGetter limits the interface of auth.ClientI to methods needed
// by NewAccessCheckerForRemoteCluster.
type CurrentUserRoleGetter interface {
// GetCurrentUserRoles returns the remote cluster roles for the current
// user, traits have not been applied.
GetCurrentUserRoles(context.Context) ([]types.Role, error)
// GetCurrentUser returns the remote cluster's view of the current user.
GetCurrentUser(context.Context) (types.User, error)
}
// NewAccessCheckerForRemoteCluster returns an AccessChecker that can check
// user's access to resources that may be located in remote/leaf Teleport
// clusters.
func NewAccessCheckerForRemoteCluster(ctx context.Context, localAccessInfo *AccessInfo, clusterName string, access CurrentUserRoleGetter) (AccessChecker, error) {
if localAccessInfo.ScopePin != nil {
return nil, trace.BadParameter("cannot create unscoped remote cluster AccessChecker based on scoped identity")
}
// Fetch the remote cluster's view of the current user's roles.
remoteRoles, err := access.GetCurrentUserRoles(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// Fetch the remote cluster's view of the current user's traits.
// These can technically be different than the local user's traits, see
// AccessInfoFromRemote(Certificate|Identity).
remoteUser, err := access.GetCurrentUser(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
remoteAccessInfo := &AccessInfo{
Username: remoteUser.GetName(),
Traits: remoteUser.GetTraits(),
// Will fill this in with the names of the remote/mapped roles we got
// from GetCurrentUserRoles.
Roles: make([]string, 0, len(remoteRoles)),
// AllowedResourceAccessIDs are always the same across clusters.
AllowedResourceAccessIDs: localAccessInfo.AllowedResourceAccessIDs,
DelegationSessionID: localAccessInfo.DelegationSessionID,
}
for i := range remoteRoles {
remoteRoles[i], err = ApplyTraitsWithContext(remoteRoles[i], RoleTemplateContext{
Username: remoteAccessInfo.Username,
Traits: remoteAccessInfo.Traits,
})
if err != nil {
return nil, trace.Wrap(err)
}
remoteAccessInfo.Roles = append(remoteAccessInfo.Roles, remoteRoles[i].GetName())
}
roleSet := NewRoleSet(remoteRoles...)
return &accessChecker{
info: remoteAccessInfo,
// localCluster is a bit of a misnomer here, but it means the local
// cluster of the resources to which access will be checked, which in
// this case may be a remote cluster. localCluster is used for access
// checks involving Resource Access Requests, the cluster name is
// included in the unique ID of the resource, the accessChecker can only
// check access to resources in that cluster.
localCluster: clusterName,
RoleSet: roleSet,
}, nil
}
type allowedResourceMatch struct {
Match *types.ResourceAccessID
}
// checkAllowedResources enforces AllowedResourceAccessIDs if present on the identity.
func (a *accessChecker) checkAllowedResources(r AccessCheckable) (allowedResourceMatch, error) {
if len(a.info.AllowedResourceAccessIDs) == 0 {
// certificate does not contain a list of specifically allowed
// resources, only role-based access control is used
return allowedResourceMatch{}, nil
}
// Note: logging in this function only happens in trace mode. This is because
// adding logging to this function (which is called on every resource returned
// by the backend) can slow down this function by 50x for large clusters!
ctx := context.Background()
isLoggingEnabled := rbacLogger.Enabled(ctx, logutils.TraceLevel)
var match *types.ResourceAccessID
for _, resourceID := range a.info.AllowedResourceAccessIDs {
id := resourceID.GetResourceID()
if id.ClusterName != a.localCluster || !matchesUCRResource(resourceID, r) {
continue
}
if resourceID.GetConstraints().Unenforceable() {
if isLoggingEnabled {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access denied, matched resource ID carries constraints this component cannot enforce",
slog.String("resource_id", types.ResourceIDToString(id)),
)
}
return allowedResourceMatch{}, trace.AccessDenied(
"access to %v %+q denied because it carries constraints this component cannot enforce; they may have been created by a newer Teleport version",
r.GetKind(), r.GetName())
}
if match == nil {
match = &resourceID
}
}
if match != nil {
// Allowed to access this resource by resource ID, move on to role checks.
if isLoggingEnabled {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Matched allowed resource ID",
slog.String("resource_id", types.ResourceIDToString(match.GetResourceID())),
)
}
return allowedResourceMatch{match}, nil
}
if isLoggingEnabled {
// We just want to log allowed IDs here; discarding additional info is ok.
allowedResources, err := types.ResourceIDsToString(types.RiskyExtractResourceIDs(a.info.AllowedResourceAccessIDs))
if err != nil {
return allowedResourceMatch{}, trace.Wrap(err)
}
slog.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, not in allowed resource IDs",
slog.String("resource_kind", r.GetKind()),
slog.String("resource_name", r.GetName()),
slog.Any("allowed_resources", allowedResources),
)
return allowedResourceMatch{}, trace.AccessDenied("access to %v denied, %q not in allowed resource IDs %s",
r.GetKind(), r.GetName(), allowedResources)
}
return allowedResourceMatch{}, trace.AccessDenied("access to %v denied, not in allowed resource IDs", r.GetKind())
}
// matchesUCRResource matches requested resource with its respective
// resource type stored in the unified resource cache.
func matchesUCRResource(requestedR types.ResourceAccessID, r AccessCheckable) bool {
if requestedR.GetResourceID().Name != r.GetName() {
return false
}
// If the allowed resource has `Kind=types.KindKubePod` or any other
// Kubernetes supported kinds - types.KubernetesResourcesKinds-, we allow the user to
// access the Kubernetes cluster that it belongs to.
// At this point, we do not verify that the accessed resource matches the
// allowed resources, but that verification happens in the caller function.
if slices.Contains(types.KubernetesResourcesKinds, requestedR.GetResourceID().Kind) || strings.HasPrefix(requestedR.GetResourceID().Kind, types.AccessRequestPrefixKindKube) {
return r.GetKind() == types.KindKubernetesCluster
}
// Identity Center account is stored as KindApp kind and
// KindIdentityCenterAccount subKind in the unified resource cache.
if requestedR.GetResourceID().Kind == types.KindIdentityCenterAccount {
return r.GetKind() == types.KindApp && r.GetSubKind() == types.KindIdentityCenterAccount
}
return requestedR.GetResourceID().Kind == r.GetKind()
}
// AccessInfo returns the AccessInfo that this access checker is based on.
func (a *accessChecker) AccessInfo() *AccessInfo {
return a.info
}
// DelegationSessionID returns the ID of the current Delegation Session.
func (a *accessChecker) DelegationSessionID() string {
return a.info.DelegationSessionID
}
// blockedInDelegationSession checks whether the given action is disallowed
// because the caller is in a Delegation Session with restricted access to
// specific resources only.
//
// Without this check, the `AllowedResourceAccessIDs` would only restrict
// regular access (e.g. SSH-ing into a node), not administrative actions,
// so if the delegating user has a role that allows them to mutate resources,
// the session user would also be able to do this on their behalf.
//
// If the Delegation Session has a "wildcard" resource selector, the user
// has explicitly allowed the session user to take on *all* of their
// permissions, including destructive administrative actions.
func (a *accessChecker) blockedInDelegationSession(kind, verb string) bool {
if a.DelegationSessionID() == "" || len(a.GetAllowedResourceAccessIDs()) == 0 {
return false
}
// Collect all the resource kinds the session has access to.
allowedKinds := set.New[string]()
for _, id := range a.GetAllowedResourceAccessIDs() {
allowedKinds.Add(id.GetResourceID().Kind)
}
// Also add the implied resource kinds (e.g. app -> app_server).
impliedKinds := map[string][]string{
types.KindApp: []string{types.KindAppServer},
types.KindDatabase: []string{types.KindDatabaseServer},
types.KindKubernetesCluster: []string{types.KindKubeServer},
types.KindWindowsDesktop: []string{types.KindWindowsDesktopService},
}
for parent, children := range impliedKinds {
if allowedKinds.Contains(parent) {
allowedKinds.Add(children...)
}
}
// These verbs are allowed to enable `tsh ls`, etc.
allowedVerbs := set.New(types.VerbList, types.VerbRead, types.VerbReadNoSecrets)
return !allowedKinds.Contains(kind) || !allowedVerbs.Contains(verb)
}
// CheckAccessToRule checks access to a rule within a namespace.
//
// It extends [RoleSet.CheckAccessToRule] to prevent Delegation Sessions with
// restricted access to specific resources from inheriting the user's destructive
// admin/rule based privileges
func (a *accessChecker) CheckAccessToRule(ctx RuleContext, namespace string, resource string, verb string) error {
if a.blockedInDelegationSession(resource, verb) {
return trace.AccessDenied("access denied to perform action %q on %q", verb, resource)
}
return a.RoleSet.CheckAccessToRule(ctx, namespace, resource, verb)
}
// GuessIfAccessIsPossible guesses if access is possible for an entire category
// of resources.
func (a *accessChecker) GuessIfAccessIsPossible(ctx RuleContext, namespace string, resource string, verb string) error {
if a.blockedInDelegationSession(resource, verb) {
return trace.AccessDenied("access denied to perform action %q on %q", verb, resource)
}
return a.RoleSet.GuessIfAccessIsPossible(ctx, namespace, resource, verb)
}
// ExtractConditionForIdentifier returns a restrictive filter expression
// for list queries based on the rules' `where` conditions.
func (a *accessChecker) ExtractConditionForIdentifier(ctx RuleContext, namespace, resource, verb, identifier string) (*types.WhereExpr, error) {
if a.blockedInDelegationSession(resource, verb) {
return nil, trace.AccessDenied("access denied to perform action %q on %q", verb, resource)
}
return a.RoleSet.ExtractConditionForIdentifier(ctx, namespace, resource, verb, identifier)
}
// CheckAccess checks if the identity for this AccessChecker has access to the given resource.
func (a *accessChecker) CheckAccess(r AccessCheckable, state AccessState, matchers ...RoleMatcher) error {
// Immediately return an error regardless of potential preconditions. This is to maintain backwards compatibility
// with existing callers of CheckAccess which expect an error when access is denied.
state.ReturnPreconditions = false
_, err := a.validateAccessConditions(r, state, matchers...)
return trace.Wrap(err)
}
// CheckConditionalAccess checks if the identity for this AccessChecker has conditional access to the given resource.
func (a *accessChecker) CheckConditionalAccess(r AccessCheckable, state AccessState, matchers ...RoleMatcher) ([]*decisionpb.Precondition, error) {
// Indicate that we want preconditions to be returned if access is granted rather than an error.
state.ReturnPreconditions = true
return a.validateAccessConditions(r, state, matchers...)
}
func (a *accessChecker) validateAccessConditions(r AccessCheckable, state AccessState, matchers ...RoleMatcher) ([]*decisionpb.Precondition, error) {
// Enforce AllowedResourceAccessIDs if present; capture match
res, err := a.checkAllowedResources(r)
if err != nil {
return nil, trace.Wrap(err)
}
switch rr := r.(type) {
case types.Resource153UnwrapperT[IdentityCenterAccount]:
matchers = append(matchers, NewIdentityCenterAccountMatcher(rr.UnwrapT()))
case types.Resource153UnwrapperT[IdentityCenterAccountAssignment]:
matchers = append(matchers, NewIdentityCenterAccountAssignmentMatcher(rr.UnwrapT()))
}
// If matched RID has ResourceConstraints, guard any principal-bearing matcher(s)
if res.Match != nil && res.Match.GetConstraints() != nil {
guard := WithConstraints(res.Match.GetConstraints())
for i := range matchers {
matchers[i] = guard(matchers[i])
}
}
preconds, err := a.checkAccess(r, a.info.Username, a.info.Traits, state, matchers...)
if err != nil {
return nil, trace.Wrap(err)
}
return preconds, nil
}
// CheckAccessToSAMLIdP checks access to SAML IdP service provider resource.
// It checks for both the legacy RBAC (role v7 and below) that checks for IDP
// role option and MFA, as well as non-legacy RBAC (role v8 and above) that checks
// for labels, MFA and Device Trust.
func (a *accessChecker) CheckAccessToSAMLIdP(r AccessCheckable, authPref readonly.AuthPreference, state AccessState, matchers ...RoleMatcher) error {
if _, err := a.checkAllowedResources(r); err != nil {
return trace.Wrap(err)
}
return trace.Wrap(a.RoleSet.CheckAccessToSAMLIdP(r, a.info.Username, a.info.Traits, authPref, state, matchers...))
}
// GetKubeResources returns the allowed and denied Kubernetes Resources configured
// for a user.
func (a *accessChecker) GetKubeResources(cluster types.KubeCluster) (allowed, denied []types.KubernetesResource) {
if len(a.info.AllowedResourceAccessIDs) == 0 {
return a.RoleSet.GetKubeResources(cluster, a.info.Username, a.info.Traits)
}
var err error
rolesAllowed, rolesDenied := a.RoleSet.GetKubeResources(cluster, a.info.Username, a.info.Traits)
// If we have a legacy 'namespace' in the allowedResourceIDs, we need to add the new 'namespaces' one.
// The old one will get mapped to wildcard later.
allowedResourceAccessIDs := slices.Clone(a.info.AllowedResourceAccessIDs)
for _, elem := range a.info.AllowedResourceAccessIDs {
if rid := elem.GetResourceID(); rid.Kind == types.KindKubeNamespace {
allowedResourceAccessIDs = append(allowedResourceAccessIDs, types.ResourceAccessID{Id: types.ResourceID{
ClusterName: rid.ClusterName,
Kind: types.AccessRequestPrefixKindKubeClusterWide + "namespaces",
SubResourceName: rid.SubResourceName,
Name: rid.Name,
}})
}
}
// Allways append the denied resources from the roles. This is because
// the denied resources from the roles take precedence over the allowed
// resources from the certificate.
denied = rolesDenied
for _, wr := range allowedResourceAccessIDs {
r := wr.GetResourceID()
if r.Name != cluster.GetName() || r.ClusterName != a.localCluster {
continue
}
switch {
case slices.Contains(types.KubernetesResourcesKinds, r.Kind) || strings.HasPrefix(r.Kind, types.AccessRequestPrefixKindKube):
namespace := ""
name := ""
if slices.Contains(types.KubernetesClusterWideResourceKinds, r.Kind) || strings.HasPrefix(r.Kind, types.AccessRequestPrefixKindKubeClusterWide) {
// Cluster wide resources do not have a namespace.
name = r.SubResourceName
r.Kind = strings.TrimPrefix(r.Kind, types.AccessRequestPrefixKindKubeClusterWide)
} else {
r.Kind = strings.TrimPrefix(r.Kind, types.AccessRequestPrefixKindKubeNamespaced)
splitted := strings.SplitN(r.SubResourceName, "/", 3)
// This condition should never happen since SubResourceName is validated
// but it's better to validate it.
if len(splitted) != 2 {
continue
}
namespace = splitted[0]
// namespace * would also include cluster-wide resources, if we
// have a wildcard with a known namespaced resource, use a pattern
// that will not match cluster-wide resources.
if namespace == types.Wildcard {
namespace = "^.+$"
}
name = splitted[1]
}
// Map legacy names to the new ones.
kind := types.KubernetesResourcesKindsPlurals[r.Kind]
if kind == "" {
kind = r.Kind
}
// NOTE: The kind 'namespace' behavior changed, to maintain backwards compatibility,
// map the legacy value to wildcard.
if r.Kind == types.KindKubeNamespace {
// When requesting the legacy "namespace" kind, we include all api groups.
kind = types.Wildcard + "." + types.Wildcard
namespace = name
// namespace * would also include cluster-wide resources, if we
// have a wildcard with the legacy "namespace" kind, use a pattern
// that will not match cluster-wide resources.
if namespace == types.Wildcard {
namespace = "^.+$"
}
name = types.Wildcard
}
gk := schema.ParseGroupKind(kind)
if gk.Group == "" {
gk.Group = types.KubernetesResourcesV7KindGroups[r.Kind]
}
r := types.KubernetesResource{
Kind: gk.Kind,
Namespace: namespace,
Name: name,
APIGroup: gk.Group,
}
// matchKubernetesResource checks if the Kubernetes Resource matches the tuple
// (kind, namespace, kame) from the allowed/denied list and does not match the resource
// verbs. Verbs are not checked here because they are not included in the
// ResourceID but we collect them and set them in the returned KubernetesResource
// so that they can be matched when the resource is accessed.
if r.Verbs, err = matchKubernetesResource(r, rolesAllowed, rolesDenied); err == nil {
allowed = append(allowed, r)
}
case r.Kind == types.KindKubernetesCluster:
// When a user has access to a Kubernetes cluster through Resource Access request,
// he has access to all resources in that cluster that he has access to through his roles.
// In that case, we append the allowed and denied resources from the roles.
return rolesAllowed, rolesDenied
}
}
return append(allowed, types.KubernetesResourceSelfSubjectAccessReview), denied
}
// matchKubernetesResource checks if the Kubernetes Resource does not match any
// entry from the deny list and matches at least one entry from the allowed list.
func matchKubernetesResource(resource types.KubernetesResource, allowed, denied []types.KubernetesResource) ([]string, error) {
// utils.KubeResourceMatchesRegex checks if the resource.Kind is strictly equal
// to each entry and validates if the Name and Namespace fields matches the
// regex allowed by each entry.
result, _, err := utils.KubeResourceMatchesRegexWithVerbsCollector(resource, denied)
if err != nil {
return nil, trace.Wrap(err)
} else if result {
return nil, trace.AccessDenied("access to %s %q denied", resource.Kind, resource.ClusterResource())
}
result, verbs, err := utils.KubeResourceMatchesRegexWithVerbsCollector(resource, allowed)
if err != nil {
return nil, trace.Wrap(err)
} else if !result {
return nil, trace.AccessDenied("access to %s %q denied", resource.Kind, resource.ClusterResource())
}
return verbs, nil
}
// GetAllowedResourceAccessIDs returns the list of allowed resources the identity for
// the AccessChecker is allowed to access. An empty or nil list indicates that
// there are no resource-specific restrictions.
func (a *accessChecker) GetAllowedResourceAccessIDs() []types.ResourceAccessID {
return a.info.AllowedResourceAccessIDs
}
// Traits returns the set of user traits
func (a *accessChecker) Traits() wrappers.Traits {
return a.info.Traits
}
// DatabaseAutoUserMode returns whether a user should be auto-created in
// the database.
func (a *accessChecker) DatabaseAutoUserMode(database types.Database) (types.CreateDatabaseUserMode, error) {
result, err := a.checkDatabaseRoles(database)
return result.createDatabaseUserMode(), trace.Wrap(err)
}
// CheckDatabaseRoles returns whether a user should be auto-created in the
// database and a list of database roles to assign.
func (a *accessChecker) CheckDatabaseRoles(database types.Database, userRequestedRoles []string) ([]string, error) {
result, err := a.checkDatabaseRoles(database)
if err != nil {
return nil, trace.Wrap(err)
}
switch {
case !result.createDatabaseUserMode().IsEnabled():
return []string{}, nil
// If user requested a list of roles, make sure all requested roles are
// allowed.
case len(userRequestedRoles) > 0:
for _, requestedRole := range userRequestedRoles {
if !slices.Contains(result.allowedRoles(), requestedRole) {
return nil, trace.AccessDenied("access to database role %q denied", requestedRole)
}
}
return userRequestedRoles, nil
// If user does not provide any roles, use all allowed roles from roleset.
default:
return result.allowedRoles(), nil
}
}
type checkDatabaseRolesResult struct {
allowedRoleSet RoleSet
deniedRoleSet RoleSet
}
func (result *checkDatabaseRolesResult) createDatabaseUserMode() types.CreateDatabaseUserMode {
if result == nil {
return types.CreateDatabaseUserMode_DB_USER_MODE_UNSPECIFIED
}
return result.allowedRoleSet.GetCreateDatabaseUserMode()
}
func (result *checkDatabaseRolesResult) allowedRoles() []string {
if result == nil {
return nil
}
rolesMap := set.New[string]()
for _, role := range result.allowedRoleSet {
for _, dbRole := range role.GetDatabaseRoles(types.Allow) {
rolesMap.Add(dbRole)
}
}
for _, role := range result.deniedRoleSet {
for _, dbRole := range role.GetDatabaseRoles(types.Deny) {
rolesMap.Remove(dbRole)
}
}
// The database user provisioning code is picky - it requires a non-nil
// slice of roles, because this value is passed directly to a SQL query.
return rolesMap.ElementsNotNil()
}
func (a *accessChecker) checkDatabaseRoles(database types.Database) (*checkDatabaseRolesResult, error) {
// First, collect roles from this roleset that have create database user mode set.
var autoCreateRoles RoleSet
for _, role := range a.RoleSet {
if role.GetCreateDatabaseUserMode().IsEnabled() {
autoCreateRoles = append(autoCreateRoles, role)
}
}
// If there are no "auto-create user" roles, nothing to do.
if len(autoCreateRoles) == 0 {
return nil, nil
}
// Otherwise, iterate over auto-create roles matching the database user
// is connecting to and compile a list of roles database user should be
// assigned.
var allowedRoleSet RoleSet
for _, role := range autoCreateRoles {
match, _, err := checkRoleLabelsMatch(types.Allow, role, a.info.Username, a.info.Traits, database, false)
if err != nil {
return nil, trace.Wrap(err)
}
if !match {
continue
}
allowedRoleSet = append(allowedRoleSet, role)
}
var deniedRoleSet RoleSet
for _, role := range autoCreateRoles {
match, _, err := checkRoleLabelsMatch(types.Deny, role, a.info.Username, a.info.Traits, database, false)
if err != nil {
return nil, trace.Wrap(err)
}
if !match {
continue
}
deniedRoleSet = append(deniedRoleSet, role)
}
// The collected role list can be empty and that should be ok, we want to
// leave the behavior of what happens when a user is created with default
// "no roles" configuration up to the target database.
result := checkDatabaseRolesResult{
allowedRoleSet: allowedRoleSet,
deniedRoleSet: deniedRoleSet,
}
return &result, nil
}
// GetDatabasePermissions returns a set of database permissions applicable for the user in the context of particular database.
func (a *accessChecker) GetDatabasePermissions(database types.Database) (allow types.DatabasePermissions, deny types.DatabasePermissions, err error) {
result, err := a.checkDatabaseRoles(database)
if err != nil {
return nil, nil, trace.Wrap(err)
}
if !result.createDatabaseUserMode().IsEnabled() {
return nil, nil, nil
}
for _, role := range result.allowedRoleSet {
allow = append(allow, role.GetDatabasePermissions(types.Allow)...)
}
for _, role := range result.deniedRoleSet {
deny = append(deny, role.GetDatabasePermissions(types.Deny)...)
}
return allow, deny, nil
}
// EnumerateDatabaseUsers specializes EnumerateEntities to enumerate db_users.
func (a *accessChecker) EnumerateDatabaseUsers(database types.Database, extraUsers ...string) (EnumerationResult, error) {
// When auto-user provisioning is enabled, only Teleport username is allowed.
if database.IsAutoUsersEnabled() {
result := NewEnumerationResult()
autoUser, err := a.DatabaseAutoUserMode(database)
if err != nil {
return result, trace.Wrap(err)
} else if autoUser.IsEnabled() {
result.allowedDeniedMap[a.info.Username] = true
return result, nil
}
}
listFn := func(role types.Role, condition types.RoleConditionType) []string {
return role.GetDatabaseUsers(condition)
}
newMatcher := func(user string) RoleMatcher {
return NewDatabaseUserMatcher(database, user)
}
return a.EnumerateEntities(database, listFn, newMatcher, extraUsers...), nil
}
// EnumerateDatabaseNames specializes EnumerateEntities to enumerate db_names.
func (a *accessChecker) EnumerateDatabaseNames(database types.Database, extraNames ...string) EnumerationResult {
listFn := func(role types.Role, condition types.RoleConditionType) []string {
return role.GetDatabaseNames(condition)
}
newMatcher := func(dbName string) RoleMatcher {
return &DatabaseNameMatcher{Name: dbName}
}
return a.EnumerateEntities(database, listFn, newMatcher, extraNames...)
}
// EnumerateMCPTools specializes EnumerateEntities to enumerate mcp.tools.
func (a *accessChecker) EnumerateMCPTools(app types.Application) EnumerationResult {
listFn := func(role types.Role, condition types.RoleConditionType) []string {
if mcpSpec := role.GetMCPPermissions(condition); mcpSpec != nil {
return mcpSpec.Tools
}
return nil
}
// Do not use MCPToolMatcher. We are enumerating the expressions.
newMatcher := func(toolRegex string) RoleMatcher {
return RoleMatcherFunc(func(role types.Role, condition types.RoleConditionType) (bool, error) {
if mcpSpec := role.GetMCPPermissions(condition); mcpSpec != nil {
return slices.Contains(mcpSpec.Tools, toolRegex), nil
}
return false, nil
})
}
return a.EnumerateEntities(app, listFn, newMatcher)
}
// roleEntitiesListFn is used for listing a role's allowed/denied entities.
type roleEntitiesListFn func(types.Role, types.RoleConditionType) []string
// roleMatcherFactoryFn is used for making a role matcher for a given entity.
type roleMatcherFactoryFn func(entity string) RoleMatcher
// EnumerateEntities works on a given role set to return a minimal description
// of allowed set of entities (db_users, db_names, etc). It is biased towards
// *allowed* entities; It is meant to describe what the user can do, rather than
// cannot do. For that reason if the user isn't allowed to pick *any* entities,
// the output will be empty.
//
// In cases where * is listed in set of allowed entities, it may be hard for
// users to figure out the expected entity to use. For this reason the parameter
// extraEntities provides an extra set of entities to be checked against
// RoleSet. This extra set of entities may be sourced e.g. from user connection
// history.
func (a *accessChecker) EnumerateEntities(resource AccessCheckable, listFn roleEntitiesListFn, newMatcher roleMatcherFactoryFn, extraEntities ...string) EnumerationResult {
result := NewEnumerationResult()
// gather entities for checking from the roles, check wildcards.
var entities []string
for _, role := range a.RoleSet {
wildcardAllowed := false
wildcardDenied := false
// Only append allowed entries and update wildcardAllowed if the role
// allows the resource without any matcher. In the real CheckAccess,
// RoleMatchers(matchers).MatchAll(role, types.Allow) is only run when
// namespace and label matching passes on this resource. Checking
// if the role allows the resource without any matcher confirms
// namespace and label matching has passed.
var resourceAllowedByRole bool
if _, err := NewRoleSet(role).checkAccess(resource, a.info.Username, a.info.Traits, AccessState{MFAVerified: true}); err == nil {
resourceAllowedByRole = true
}
for _, e := range listFn(role, types.Allow) {
if e == types.Wildcard {
wildcardAllowed = true
} else if resourceAllowedByRole {
entities = append(entities, e)
}
}
for _, e := range listFn(role, types.Deny) {
if e == types.Wildcard {
wildcardDenied = true
} else {
entities = append(entities, e)
}
}
result.wildcardDenied = result.wildcardDenied || wildcardDenied
if resourceAllowedByRole {
result.wildcardAllowed = result.wildcardAllowed || wildcardAllowed
}
}
entities = apiutils.Deduplicate(append(entities, extraEntities...))
// check each individual role spec entity against the resource.
for _, e := range entities {
err := a.CheckAccess(resource, AccessState{MFAVerified: true}, newMatcher(e))
result.allowedDeniedMap[e] = err == nil
}
return result
}
// GetAllowedLoginsForResource returns all of the allowed logins for the passed resource.
//
// Supports the following resource types:
//
// - types.Server with GetKind() == types.KindNode
// - types.KindWindowsDesktop
// - types.KindApp with IsAWSConsole() == true
func (a *accessChecker) GetAllowedLoginsForResource(resource AccessCheckable) ([]string, error) {
// Create a map indexed by all logins in the RoleSet,
// mapped to false if any role has it in its deny section,
// true otherwise.
mapped := make(map[string]bool)
resourceAsApp, resourceIsApp := resource.(interface{ IsAWSConsole() bool })
for _, role := range a.RoleSet {
var loginGetter func(types.RoleConditionType) []string
switch resource.GetKind() {
case types.KindNode:
loginGetter = role.GetLogins
case types.KindWindowsDesktop:
loginGetter = role.GetWindowsLogins
case types.KindLinuxDesktop:
loginGetter = role.GetLinuxDesktopLogins
case types.KindApp:
if !resourceIsApp {
return nil, trace.BadParameter("received unsupported resource type for Application kind: %T", resource)
}
// For Apps, only AWS currently supports listing the possible logins.
if !resourceAsApp.IsAWSConsole() {
return nil, nil
}
loginGetter = role.GetAWSRoleARNs
default:
return nil, trace.BadParameter("received unsupported resource kind: %s", resource.GetKind())
}
for _, login := range loginGetter(types.Allow) {
// Only set to true if not already set, the login is denied if any
// role denies it.
if _, alreadySet := mapped[login]; !alreadySet {
mapped[login] = true
}
}
for _, login := range loginGetter(types.Deny) {
mapped[login] = false
}
}
// Create a list of only the logins not denied by a role in the set.
var notDenied []string
for login, isNotDenied := range mapped {
if isNotDenied {
notDenied = append(notDenied, login)
}
}
var newLoginMatcher func(login string) RoleMatcher
switch resource.GetKind() {
case types.KindNode:
newLoginMatcher = NewLoginMatcher
case types.KindWindowsDesktop:
newLoginMatcher = NewWindowsLoginMatcher
case types.KindLinuxDesktop:
newLoginMatcher = NewLinuxDesktopLoginMatcher
case types.KindApp:
if !resourceIsApp || !resourceAsApp.IsAWSConsole() {
return nil, trace.BadParameter("received unsupported resource type for Application: %T", resource)
}
newLoginMatcher = NewAppAWSLoginMatcher
default:
return nil, trace.BadParameter("received unsupported resource kind: %s", resource.GetKind())
}
// Filter the not-denied logins for those allowed to be used with the given resource.
var allowed []string
for _, login := range notDenied {
err := a.CheckAccess(resource, AccessState{MFAVerified: true}, newLoginMatcher(login))
if err == nil {
allowed = append(allowed, login)
}
}
return allowed, nil
}
// CheckAccessToRemoteCluster checks if a role has access to remote cluster. Deny rules are
// checked first then allow rules. Access to a cluster is determined by
// namespaces, labels, and logins.
func (a *accessChecker) CheckAccessToRemoteCluster(rc types.RemoteCluster) error {
if len(a.RoleSet) == 0 {
return trace.AccessDenied("access to cluster denied")
}
// Note: logging in this function only happens in trace mode, this is because
// adding logging to this function (which is called on every server returned
// by GetRemoteClusters) can slow down this function by 50x for large clusters!
ctx := context.Background()
isLoggingEnabled := rbacLogger.Enabled(ctx, logutils.TraceLevel)
rcLabels := rc.GetMetadata().Labels
// For backwards compatibility, if there is no role in the set with label
// matchers and the cluster has no labels, assume that the role set has
// access to the cluster.
usesLabels := false
for _, role := range a.RoleSet {
unset, err := labelMatchersUnset(role, types.KindRemoteCluster)
if err != nil {
return trace.Wrap(err)
}
if !unset {
usesLabels = true
break
}
}
if !usesLabels && len(rcLabels) == 0 {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Grant access to cluster - no role uses cluster labels and the cluster is not labeled",
slog.String("cluster_name", rc.GetName()),
slog.Any("roles", a.RoleNames()),
)
return nil
}
// Check deny rules first: a single matching label from
// the deny role set prohibits access.
var errs []error
for _, role := range a.RoleSet {
matchLabels, labelsMessage, err := checkRoleLabelsMatch(types.Deny, role, a.info.Username, a.info.Traits, rc, isLoggingEnabled)
if err != nil {
return trace.Wrap(err)
}
if matchLabels {
// This condition avoids formatting calls on large scale.
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access to cluster denied, deny rule matched",
slog.String("cluster", rc.GetName()),
slog.String("role", role.GetName()),
slog.String("label_message", labelsMessage),
)
return trace.AccessDenied("access to cluster denied")
}
}
// Check allow rules: label has to match in any role in the role set to be granted access.
for _, role := range a.RoleSet {
matchLabels, labelsMessage, err := checkRoleLabelsMatch(types.Allow, role, a.info.Username, a.info.Traits, rc, isLoggingEnabled)
if err != nil {
return trace.Wrap(err)
}
labelMatchers, err := role.GetLabelMatchers(types.Allow, types.KindRemoteCluster)
if err != nil {
return trace.Wrap(err)
}
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Check access to role",
slog.String("role", role.GetName()),
slog.String("cluster", rc.GetName()),
slog.Any("cluster_labels", rcLabels),
slog.Any("match_labels", matchLabels),
slog.String("labels_message", labelsMessage),
slog.Any("error", err),
slog.Any("allow", labelMatchers),
)
if matchLabels {
return nil
}
if isLoggingEnabled {
deniedError := trace.AccessDenied("role=%v, match(%s)",
role.GetName(), labelsMessage)
errs = append(errs, deniedError)
}
}
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access to cluster denied, no allow rule matched",
slog.String("cluster", rc.GetName()),
slog.Any("error", errs),
)
return trace.AccessDenied("access to cluster denied")
}
// DesktopGroups returns the desktop groups a user is allowed to create or an access denied error if a role disallows desktop user creation
func (a *accessChecker) DesktopGroups(s types.WindowsDesktop) ([]string, error) {
groups := set.New[string]()
for _, role := range a.RoleSet {
result, _, err := checkRoleLabelsMatch(types.Allow, role, a.info.Username, a.info.Traits, s, false)
if err != nil {
return nil, trace.Wrap(err)
}
// skip nodes that dont have matching labels
if !result {
continue
}
createDesktopUser := role.GetOptions().CreateDesktopUser
// if any of the matching roles do not enable create host
// user, the user should not be allowed on
if createDesktopUser == nil || !createDesktopUser.Value {
return nil, trace.AccessDenied("user is not allowed to create host users")
}
for _, group := range role.GetDesktopGroups(types.Allow) {
groups.Add(group)
}
}
for _, role := range a.RoleSet {
result, _, err := checkRoleLabelsMatch(types.Deny, role, a.info.Username, a.info.Traits, s, false)
if err != nil {
return nil, trace.Wrap(err)
}
if !result {
continue
}
for _, group := range role.GetDesktopGroups(types.Deny) {
groups.Remove(group)
}
}
// These groups get encoded into a certificate that's parsed by
// Rust code on Windows. That code expects an empty JSON array,
// not a null value.
return groups.ElementsNotNil(), nil
}
func convertHostUserMode(mode types.CreateHostUserMode) decisionpb.HostUserMode {
switch mode {
case types.CreateHostUserMode_HOST_USER_MODE_KEEP:
return decisionpb.HostUserMode_HOST_USER_MODE_KEEP
case types.CreateHostUserMode_HOST_USER_MODE_INSECURE_DROP:
return decisionpb.HostUserMode_HOST_USER_MODE_DROP
default:
return decisionpb.HostUserMode_HOST_USER_MODE_UNSPECIFIED
}
}
// HostUsersDecision is a decision to allow or disallow host user creation.
type HostUsersDecision struct {
// Info is host users information. If host users creation is disallowed, this will be nil.
Info *decisionpb.HostUsersInfo
// AllowedBy is a list of determinants that allow host user creation.
AllowedBy []*decisionpb.Determinant
// DeniedBy is a list of determinants that disallow host user creation.
DeniedBy []*decisionpb.Determinant
}
// HostUsers returns host user decision matching a server.
func (a *accessChecker) HostUsers(s types.Server) (*HostUsersDecision, error) {
groups := set.New[string]()
shellToRoles := make(map[string][]string)
var shell string
var mode types.CreateHostUserMode
decision := new(HostUsersDecision)
for _, role := range a.RoleSet {
result, _, err := checkRoleLabelsMatch(types.Allow, role, a.info.Username, a.info.Traits, s, false)
if err != nil {
return nil, trace.Wrap(err)
}
// skip roles that don't have matching labels
if !result {
continue
}
createHostUserMode := role.GetOptions().CreateHostUserMode
//nolint:staticcheck // this field is preserved for existing deployments, but shouldn't be used going forward
createHostUser := role.GetOptions().CreateHostUser
if createHostUserMode == types.CreateHostUserMode_HOST_USER_MODE_UNSPECIFIED {
createHostUserMode = types.CreateHostUserMode_HOST_USER_MODE_OFF
if createHostUser != nil && createHostUser.Value {
createHostUserMode = types.CreateHostUserMode_HOST_USER_MODE_KEEP
}
}
if createHostUserMode == types.CreateHostUserMode_HOST_USER_MODE_OFF {
decision.DeniedBy = append(decision.DeniedBy, decisionpb.Determinant_builder{
Kind: role.GetKind(),
Name: role.GetName(),
}.Build())
continue
}
decision.AllowedBy = append(decision.AllowedBy, decisionpb.Determinant_builder{
Kind: role.GetKind(),
Name: role.GetName(),
}.Build())
if mode == types.CreateHostUserMode_HOST_USER_MODE_UNSPECIFIED {
mode = createHostUserMode
}
// prefer to use HostUserModeKeep over InsecureDrop if mode has already been set.
if mode == types.CreateHostUserMode_HOST_USER_MODE_INSECURE_DROP &&
createHostUserMode == types.CreateHostUserMode_HOST_USER_MODE_KEEP {
mode = types.CreateHostUserMode_HOST_USER_MODE_KEEP
}
hostUserShell := role.GetOptions().CreateHostUserDefaultShell
shell = cmp.Or(shell, hostUserShell)
if hostUserShell != "" {
shellToRoles[hostUserShell] = append(shellToRoles[hostUserShell], role.GetName())
}
for _, group := range role.GetHostGroups(types.Allow) {
groups.Add(group)
}
}
// if any of the matching roles do not enable create host user, the user should not be allowed on
// Represent denial by returning the decision with a nil info field.
if len(decision.DeniedBy) > 0 {
decision.Info = nil
return decision, nil
}
if len(shellToRoles) > 1 {
b := &strings.Builder{}
for shell, roles := range shellToRoles {
fmt.Fprintf(b, "%s=%v ", shell, roles)
}
slog.WarnContext(context.Background(), "Host user shell resolution is ambiguous due to conflicting roles, consider unifying roles around a single shell",
"selected_shell", shell,
"shell_assignments", b,
)
}
for _, role := range a.RoleSet {
result, _, err := checkRoleLabelsMatch(types.Deny, role, a.info.Username, a.info.Traits, s, false)
if err != nil {
return nil, trace.Wrap(err)
}
if !result {
continue
}
for _, group := range role.GetHostGroups(types.Deny) {
groups.Remove(group)
}
}
traits := a.Traits()
var gid string
gidL := traits[constants.TraitHostUserGID]
if len(gidL) >= 1 {
gid = gidL[0]
}
var uid string
uidL := traits[constants.TraitHostUserUID]
if len(uidL) >= 1 {
uid = uidL[0]
}
decision.Info = decisionpb.HostUsersInfo_builder{
Groups: groups.Elements(),
Mode: convertHostUserMode(mode),
Uid: uid,
Gid: gid,
Shell: shell,
}.Build()
return decision, nil
}
// HostSudoers returns host sudoers entries matching a server
func (a *accessChecker) HostSudoers(s types.Server) ([]string, error) {
var sudoers []string
roleSet := slices.Clone(a.RoleSet)
slices.SortFunc(roleSet, func(a types.Role, b types.Role) int {
return strings.Compare(a.GetName(), b.GetName())
})
seenSudoers := make(map[string]struct{})
for _, role := range roleSet {
result, _, err := checkRoleLabelsMatch(types.Allow, role, a.info.Username, a.info.Traits, s, false)
if err != nil {
return nil, trace.Wrap(err)
}
// skip nodes that dont have matching labels
if !result {
continue
}
for _, sudoer := range role.GetHostSudoers(types.Allow) {
if _, ok := seenSudoers[sudoer]; ok {
continue
}
seenSudoers[sudoer] = struct{}{}
sudoers = append(sudoers, sudoer)
}
}
var finalSudoers []string
for _, role := range roleSet {
result, _, err := checkRoleLabelsMatch(types.Deny, role, a.info.Username, a.info.Traits, s, false)
if err != nil {
return nil, trace.Wrap(err)
}
if !result {
continue
}
outer:
for _, sudoer := range sudoers {
for _, deniedSudoer := range role.GetHostSudoers(types.Deny) {
if deniedSudoer == "*" {
finalSudoers = nil
break outer
}
if sudoer != deniedSudoer {
finalSudoers = append(finalSudoers, sudoer)
}
}
}
sudoers = finalSudoers
}
return sudoers, nil
}
// AccessInfoFromLocalSSHIdentity returns a new AccessInfo populated from the
// given sshca.Identity. Should only be used for cluster local users as roles
// will not be mapped.
func AccessInfoFromLocalSSHIdentity(ident *sshca.Identity) *AccessInfo {
return &AccessInfo{
Username: ident.Username,
ScopePin: ident.ScopePin,
Roles: ident.Roles,
Traits: ident.Traits,
AllowedResourceAccessIDs: ident.AllowedResourceAccessIDs,
DelegationSessionID: ident.DelegationSessionID,
}
}
// AccessInfoFromRemoteSSHIdentity returns a new AccessInfo populated from the
// given remote cluster user's ssh identity. Remote roles will be mapped to
// local roles based on the given roleMap.
func AccessInfoFromRemoteSSHIdentity(unmappedIdentity *sshca.Identity, roleMap types.RoleMap) (*AccessInfo, error) {
if unmappedIdentity.ScopePin != nil {
return nil, trace.BadParameter("scope pinning is not supported for remote SSH identities")
}
// make a shallow copy of traits to avoid modifying the original
// (don't use maps.Clone, as we want to ensure the result is an empty, but not nil, map)
traits := make(map[string][]string, len(unmappedIdentity.Traits)+1)
maps.Copy(traits, unmappedIdentity.Traits)
// Prior to Teleport 6.2 the only trait passed to the remote cluster
// was the "logins" trait set to the SSH certificate principals.
//
// Keep backwards-compatible behavior and set it in addition to the
// traits extracted from the certificate.
traits[constants.TraitLogins] = unmappedIdentity.Principals
roles, err := MapRoles(roleMap, unmappedIdentity.Roles)
if err != nil {
return nil, trace.AccessDenied("failed to map roles for user with remote roles %v: %v", unmappedIdentity.Roles, err)
}
if len(roles) == 0 {
return nil, trace.AccessDenied("no roles mapped for user with remote roles %v", unmappedIdentity.Roles)
}
slog.DebugContext(context.Background(), "Mapped remote roles to local roles and traits",
"remote_roles", unmappedIdentity.Roles,
"local_roles", roles,
"traits", traits,
)
return &AccessInfo{
Username: unmappedIdentity.Username,
Roles: roles,
Traits: traits,
AllowedResourceAccessIDs: unmappedIdentity.AllowedResourceAccessIDs,
DelegationSessionID: unmappedIdentity.DelegationSessionID,
}, nil
}
// AccessInfoFromLocalTLSIdentity returns a new AccessInfo populated from the given
// tlsca.Identity. Should only be used for cluster local users as roles will not
// be mapped.
func AccessInfoFromLocalTLSIdentity(identity tlsca.Identity) (*AccessInfo, error) {
if len(identity.Groups) == 0 && identity.ScopePin == nil {
return nil, trace.BadParameter("tls identity %q has no roles or scope pin, this may indicate a malformed certificate or one that was issued by an incompatible teleport version", identity.Username)
}
return &AccessInfo{
Username: identity.Username,
ScopePin: identity.ScopePin,
Roles: identity.Groups,
Traits: identity.Traits,
AllowedResourceAccessIDs: identity.AllowedResourceAccessIDs,
DelegationSessionID: identity.DelegationSessionID,
}, nil
}
// AccessInfoFromRemoteTLSIdentity returns a new AccessInfo populated from the
// given remote cluster user's tlsca.Identity. Remote roles will be mapped to
// local roles based on the given roleMap.
func AccessInfoFromRemoteTLSIdentity(identity tlsca.Identity, roleMap types.RoleMap) (*AccessInfo, error) {
if identity.ScopePin != nil {
return nil, trace.BadParameter("scope pinning is not supported for remote TLS identities")
}
// Set internal traits for the remote user. This allows Teleport to work by
// passing exact logins, Kubernetes users/groups, database users/names, and
// AWS Role ARNs to the remote cluster.
traits := map[string][]string{
constants.TraitLogins: identity.Principals,
constants.TraitKubeGroups: identity.KubernetesGroups,
constants.TraitKubeUsers: identity.KubernetesUsers,
constants.TraitDBNames: identity.DatabaseNames,
constants.TraitDBUsers: identity.DatabaseUsers,
constants.TraitAWSRoleARNs: identity.AWSRoleARNs,
}
// Prior to Teleport 6.2 no user traits were passed to remote clusters
// except for the internal ones specified above.
//
// To preserve backwards compatible behavior, when applying traits from user
// identity, make sure to filter out those already present in the map above.
//
// This ensures that if e.g. there's a "logins" trait in the root user's
// identity, it won't overwrite the internal "logins" trait set above
// causing behavior change.
for k, v := range identity.Traits {
if _, ok := traits[k]; !ok {
traits[k] = v
}
}
unmappedRoles := identity.Groups
roles, err := MapRoles(roleMap, unmappedRoles)
if err != nil {
return nil, trace.AccessDenied("failed to map roles for remote user %q from cluster %q with remote roles %v: %v", identity.Username, identity.TeleportCluster, unmappedRoles, err)
}
if len(roles) == 0 {
return nil, trace.AccessDenied("no roles mapped for remote user %q from cluster %q with remote roles %v", identity.Username, identity.TeleportCluster, unmappedRoles)
}
slog.DebugContext(context.Background(), "Mapped roles of remote user to local roles and traits",
"remote_roles", unmappedRoles,
"user", identity.Username,
"local_roles", roles,
"traits", traits,
)
return &AccessInfo{
Username: identity.Username,
Roles: roles,
Traits: traits,
AllowedResourceAccessIDs: identity.AllowedResourceAccessIDs,
DelegationSessionID: identity.DelegationSessionID,
}, nil
}
// UserAccessState is a representation of a user's current state required for calculating access.
type UserAccessState interface {
// GetName returns the username associated with the user state.
GetName() string
// GetRoles returns the roles associated with the user's current state.
GetRoles() []string
// GetTraits returns the traits associated with the user's current sate.
GetTraits() map[string][]string
}
// UserState is a representation of a user's current state.
type UserState interface {
UserAccessState
// GetUserType returns the user type for the user login state.
GetUserType() types.UserType
// GetLabel fetches the given user label.
GetLabel(key string) (value string, ok bool)
// IsBot returns true if the user belongs to a bot.
IsBot() bool
// GetGithubIdentities returns a list of connected GitHub identities
GetGithubIdentities() []types.ExternalIdentity
// SetGithubIdentities sets the list of connected GitHub identities
SetGithubIdentities(identities []types.ExternalIdentity)
}
// AccessInfoFromUserState return a new AccessInfo populated from the roles and
// traits held be the given user state. This should only be used in cases where the
// user does not have any active access requests (initial web login, initial
// tbot certs, tests).
func AccessInfoFromUserState(user UserAccessState) *AccessInfo {
return accessInfoFromUserState(user, user.GetRoles(), nil)
}
// ScopePinnedAccessInfoFromUserState returns a new AccessInfo populated from the
// traits held by the user and the provided scope pin. Population/verification of the
// scope pin must be performed prior to calling this function.
func ScopePinnedAccessInfoFromUserState(user UserAccessState, pin *scopesv1.Pin) *AccessInfo {
return accessInfoFromUserState(user, nil, pin)
}
func accessInfoFromUserState(user UserAccessState, roles []string, pin *scopesv1.Pin) *AccessInfo {
return &AccessInfo{
Username: user.GetName(),
Roles: roles,
ScopePin: pin,
Traits: user.GetTraits(),
}
}
/*
* 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/trace"
accessgraphsecretspb "github.com/gravitational/teleport/api/gen/proto/go/teleport/accessgraph/v1"
"github.com/gravitational/teleport/api/types/accessgraph"
)
// AccessGraphSecretsGetter is an interface for getting access graph secrets.
type AccessGraphSecretsGetter interface {
// ListAllAuthorizedKeys lists all authorized keys stored in the backend.
ListAllAuthorizedKeys(ctx context.Context, pageSize int, pageToken string) ([]*accessgraphsecretspb.AuthorizedKey, string, error)
// ListAuthorizedKeysForServer lists all authorized keys for a given hostID.
ListAuthorizedKeysForServer(ctx context.Context, hostID string, pageSize int, pageToken string) ([]*accessgraphsecretspb.AuthorizedKey, string, error)
// ListAllPrivateKeys lists all private keys stored in the backend.
ListAllPrivateKeys(ctx context.Context, pageSize int, pageToken string) ([]*accessgraphsecretspb.PrivateKey, string, error)
// ListPrivateKeysForDevice lists all private keys for a given deviceID.
ListPrivateKeysForDevice(ctx context.Context, deviceID string, pageSize int, pageToken string) ([]*accessgraphsecretspb.PrivateKey, string, error)
}
// MarshalAccessGraphAuthorizedKey marshals a [accessgraphsecretspb.AuthorizedKey] resource to JSON.
func MarshalAccessGraphAuthorizedKey(in *accessgraphsecretspb.AuthorizedKey, opts ...MarshalOption) ([]byte, error) {
if err := accessgraph.ValidateAuthorizedKey(in); err != nil {
return nil, trace.Wrap(err)
}
return MarshalProtoResource(in, opts...)
}
// UnmarshalAccessGraphAuthorizedKey unmarshals a [accessgraphsecretspb.AuthorizedKey] resource from JSON.
func UnmarshalAccessGraphAuthorizedKey(data []byte, opts ...MarshalOption) (*accessgraphsecretspb.AuthorizedKey, error) {
out, err := UnmarshalProtoResource[*accessgraphsecretspb.AuthorizedKey](data, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
if err := accessgraph.ValidateAuthorizedKey(out); err != nil {
return nil, trace.Wrap(err)
}
return out, nil
}
// MarshalAccessGraphPrivateKey marshals a [accessgraphsecretspb.PrivateKey] resource to JSON.
func MarshalAccessGraphPrivateKey(in *accessgraphsecretspb.PrivateKey, opts ...MarshalOption) ([]byte, error) {
if err := accessgraph.ValidatePrivateKey(in); err != nil {
return nil, trace.Wrap(err)
}
return MarshalProtoResource(in, opts...)
}
// UnmarshalAccessGraphPrivateKey unmarshals a [accessgraphsecretspb.PrivateKey] resource from JSON.
func UnmarshalAccessGraphPrivateKey(data []byte, opts ...MarshalOption) (*accessgraphsecretspb.PrivateKey, error) {
out, err := UnmarshalProtoResource[*accessgraphsecretspb.PrivateKey](data, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
if err := accessgraph.ValidatePrivateKey(out); err != nil {
return nil, trace.Wrap(err)
}
return out, nil
}
/*
* Teleport
* Copyright (C) 2024 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 (
"github.com/gravitational/trace"
clusterconfigpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/clusterconfig/v1"
)
// UnmarshalAccessGraphSettings unmarshals the AccessGraphSettings resource from JSON.
func UnmarshalAccessGraphSettings(data []byte, opts ...MarshalOption) (*clusterconfigpb.AccessGraphSettings, error) {
out, err := UnmarshalProtoResource[*clusterconfigpb.AccessGraphSettings](data, opts...)
return out, trace.Wrap(err)
}
// MarshalAccessGraphSettings marshals the AccessGraphSettings resource to JSON.
func MarshalAccessGraphSettings(c *clusterconfigpb.AccessGraphSettings, opts ...MarshalOption) ([]byte, error) {
bytes, err := MarshalProtoResource(c, opts...)
return bytes, trace.Wrap(err)
}
/*
* 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"
"encoding/base32"
"slices"
"strings"
"time"
"github.com/charlievieth/strcase"
"github.com/gravitational/trace"
"golang.org/x/text/cases"
accesslistclient "github.com/gravitational/teleport/api/client/accesslist"
accesslistv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/accesslist/v1"
"github.com/gravitational/teleport/api/types/accesslist"
"github.com/gravitational/teleport/lib/accesslists"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils"
)
var _ AccessLists = (*accesslistclient.Client)(nil)
// AccessListsGetter defines an interface for reading access lists.
type AccessListsGetter interface {
AccessListMembersGetter
// GetAccessLists returns a list of all access lists.
GetAccessLists(context.Context) ([]*accesslist.AccessList, error)
// ListAccessLists returns a paginated list of access lists.
ListAccessLists(context.Context, int, string) ([]*accesslist.AccessList, string, error)
// ListAccessListsV2 returns a filtered and sorted paginated list of access lists.
ListAccessListsV2(context.Context, *accesslistv1.ListAccessListsV2Request) ([]*accesslist.AccessList, string, error)
// GetAccessListsToReview returns access lists that the user needs to review.
GetAccessListsToReview(context.Context) ([]*accesslist.AccessList, error)
// GetInheritedGrants returns grants inherited by access list accessListID from parent access lists.
GetInheritedGrants(context.Context, string) (*accesslist.Grants, error)
// ListUserAccessLists returns a paginated list of all access lists where the
// user is explicitly an owner or member.
ListUserAccessLists(context.Context, *accesslistv1.ListUserAccessListsRequest) ([]*accesslist.AccessList, string, error)
}
// AccessListsSuggestionsGetter defines an interface for reading access lists suggestions.
type AccessListsSuggestionsGetter interface {
// GetSuggestedAccessLists returns a list of access lists that are suggested for a given request.
GetSuggestedAccessLists(ctx context.Context, accessRequestID string) ([]*accesslist.AccessList, error)
}
// AccessLists defines an interface for managing AccessLists.
type AccessLists interface {
AccessListsGetter
AccessListsSuggestionsGetter
AccessListMembers
AccessListReviews
// UpsertAccessList creates or updates an access list resource.
UpsertAccessList(context.Context, *accesslist.AccessList) (*accesslist.AccessList, error)
// UpdateAccessList updates an access list resource.
UpdateAccessList(context.Context, *accesslist.AccessList) (*accesslist.AccessList, error)
// DeleteAccessList removes the specified access list resource.
DeleteAccessList(context.Context, string) error
// DeleteAccessList removes the specified access list resource.
DeleteAccessListV2(context.Context, *accesslistv1.DeleteAccessListRequest) error
// UpsertAccessListWithMembers creates or updates an access list resource and its members.
UpsertAccessListWithMembers(context.Context, *accesslist.AccessList, []*accesslist.AccessListMember) (*accesslist.AccessList, []*accesslist.AccessListMember, error)
// AccessRequestPromote promotes an access request to an access list.
AccessRequestPromote(ctx context.Context, req *accesslistv1.AccessRequestPromoteRequest) (*accesslistv1.AccessRequestPromoteResponse, error)
}
// AccessListsInternal extends the public AccessList interface with internal-only
// methods.
type AccessListsInternal interface {
AccessLists
// UpdateAccessListAndOverwriteMembers conditionally updates the access list,
// overwriting the list's members if successful.
UpdateAccessListAndOverwriteMembers(context.Context, *accesslist.AccessList, []*accesslist.AccessListMember) (*accesslist.AccessList, []*accesslist.AccessListMember, error)
// CleanupAccessListStatus removes invalid Status.OwnerOf and Status.MemberOf references.
CleanupAccessListStatus(ctx context.Context, accessListName string) (*accesslist.AccessList, error)
// CleanupAccessListStatusV2 removes invalid Status.(Scoped)OwnerOf and Status.(Scoped)MemberOf references.
CleanupAccessListStatusV2(ctx context.Context, accessListName accesslists.NormalizedSQN) (*accesslist.AccessList, error)
// EnsureNestedAccessListStatuses goes over all nested owners and nested members of the named
// access list and ensures nested lists' statuses owner_of/member_of contain the access list name.
EnsureNestedAccessListStatuses(ctx context.Context, accessListName string) error
// EnsureNestedAccessListStatusesV2 goes over all nested owners and nested members of the named
// access list and ensures nested lists' statuses (scoped_)owner_of/(scoped_)member_of contain
// the access list name.
EnsureNestedAccessListStatusesV2(ctx context.Context, accessListName accesslists.NormalizedSQN) error
// InsertAccessListCollection inserts a complete collection of access lists and their members from a single
// upstream source (e.g. EntraID) using a batch operation for improved performance.
//
// This method is designed for bulk import scenarios where an entire access list collection needs to be
// synchronized from an external source. All access lists and members in the collection are
// inserted using chunked batch operations, minimizing memory allocation while still reducing
// the number of write operations. Due to the batch nature of this operation (access list hierarchy
// is known upfront), we can avoid per-access-list locking and global locks to improve performance.
//
// Important: This method assumes the collection is self-contained. Access lists in the collection
// cannot reference access lists outside the collection as members or owners. This is intentional for
// collections representing a complete snapshot from a single upstream source.
// The function should be used only once during initial import where
// we are sure that Teleport doesn't have any pre-existing access lists from the upstream and the
// internal relation between upstream access lists and internal access lists doesn't exist yet.
//
// Operation can fail due to backend shutdown. In that case, if partial state was created,
// use UpsertAccessListWithMembers/DeleteAccessListMember to reconcile to the desired state.
InsertAccessListCollection(ctx context.Context, collection *accesslists.Collection) error
}
// MarshalAccessList marshals the access list resource to JSON.
func MarshalAccessList(accessList *accesslist.AccessList, opts ...MarshalOption) ([]byte, error) {
if err := accessList.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *accessList
copy.SetRevision("")
accessList = ©
}
return utils.FastMarshal(accessList)
}
// UnmarshalAccessList unmarshals the access list resource from JSON.
func UnmarshalAccessList(data []byte, opts ...MarshalOption) (*accesslist.AccessList, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing access list data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var accessList accesslist.AccessList
if err := utils.FastUnmarshal(data, &accessList); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := accessList.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
accessList.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
accessList.SetExpiry(cfg.Expires)
}
return &accessList, nil
}
// ImplicitAccessListError indicates that an operation that only makes sense for
// AccessLists with an explicit Member list has been attempted on an implicit-
// membership AccessList
type ImplicitAccessListError struct{}
// Error implements the `error` interface for ImplicitAccessListError
func (ImplicitAccessListError) Error() string {
return "requested AccessList does not have explicit member list"
}
// AccessListMemberGetter defines an interface that can retrieve access list members.
type AccessListMemberGetter interface {
// GetAccessListMember returns the specified access list member resource.
GetAccessListMember(ctx context.Context, accessList string, memberName string) (*accesslist.AccessListMember, error)
// GetAccessListMemberV2 returns the specified access list member resource.
GetAccessListMemberV2(ctx context.Context, req *accesslistv1.GetAccessListMemberRequest) (*accesslist.AccessListMember, error)
// GetAccessList returns the specified access list resource.
GetAccessList(context.Context, string) (*accesslist.AccessList, error)
// GetAccessListV2 returns the specified access list resource.
GetAccessListV2(ctx context.Context, req *accesslistv1.GetAccessListRequest) (*accesslist.AccessList, error)
// GetAccessLists returns a list of all access lists.
GetAccessLists(context.Context) ([]*accesslist.AccessList, error)
}
// AccessListMembersGetter defines an interface for reading access list members.
type AccessListMembersGetter interface {
AccessListMemberGetter
// CountAccessListMembers will count all access list members.
CountAccessListMembers(ctx context.Context, accessListName string) (membersCount uint32, listCount uint32, err error)
// CountAccessListMembersV2 will count all access list members.
CountAccessListMembersV2(ctx context.Context, req *accesslistv1.CountAccessListMembersRequest) (membersCount uint32, listCount uint32, err error)
// ListAccessListMembers returns a paginated list of all access list members.
ListAccessListMembers(ctx context.Context, accessListName string, pageSize int, pageToken string) (members []*accesslist.AccessListMember, nextToken string, err error)
// ListAccessListMembersV2 returns a paginated list of all access list members.
ListAccessListMembersV2(ctx context.Context, req *accesslistv1.ListAccessListMembersRequest) (members []*accesslist.AccessListMember, nextToken string, err error)
// ListAllAccessListMembers returns a paginated list of all access list members for all access lists.
ListAllAccessListMembers(ctx context.Context, pageSize int, pageToken string) (members []*accesslist.AccessListMember, nextToken string, err error)
// ListAllAccessListMembers returns a paginated list of all access list members for all access lists.
ListAllAccessListMembersV2(ctx context.Context, req *accesslistv1.ListAllAccessListMembersRequest) (members []*accesslist.AccessListMember, nextToken string, err error)
// GetAccessListOwners returns a list of all owners in an Access List, including those inherited from nested Access Lists.
GetAccessListOwners(ctx context.Context, accessList string) ([]*accesslist.Owner, error)
// GetAccessListOwnersV2 returns a list of all owners in an Access List, including those inherited from nested Access Lists.
GetAccessListOwnersV2(ctx context.Context, req *accesslistv1.GetAccessListOwnersRequest) ([]*accesslist.Owner, error)
}
// AccessListMembers defines an interface for managing AccessListMembers.
type AccessListMembers interface {
AccessListMembersGetter
// UpsertAccessListMember creates or updates an access list member resource.
UpsertAccessListMember(ctx context.Context, member *accesslist.AccessListMember) (*accesslist.AccessListMember, error)
// UpdateAccessListMember conditionally updates an access list member resource.
UpdateAccessListMember(ctx context.Context, member *accesslist.AccessListMember) (*accesslist.AccessListMember, error)
// DeleteAccessListMember hard deletes the specified access list member resource.
DeleteAccessListMember(ctx context.Context, accessList string, memberName string) error
// DeleteAccessListMemberV2 hard deletes the specified access list member resource.
DeleteAccessListMemberV2(ctx context.Context, req *accesslistv1.DeleteAccessListMemberRequest) error
// DeleteAllAccessListMembersForAccessList hard deletes all access list members for an access list.
DeleteAllAccessListMembersForAccessList(ctx context.Context, accessList string) error
// DeleteAllAccessListMembersForAccessListV2 hard deletes all access list members for an access list.
DeleteAllAccessListMembersForAccessListV2(ctx context.Context, req *accesslistv1.DeleteAllAccessListMembersForAccessListRequest) error
}
// MarshalAccessListMember marshals the access list member resource to JSON.
func MarshalAccessListMember(member *accesslist.AccessListMember, opts ...MarshalOption) ([]byte, error) {
if err := member.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *member
copy.SetRevision("")
member = ©
}
return utils.FastMarshal(member)
}
// UnmarshalAccessListMember unmarshals the access list member resource from JSON.
func UnmarshalAccessListMember(data []byte, opts ...MarshalOption) (*accesslist.AccessListMember, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing access list member data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var member accesslist.AccessListMember
if err := utils.FastUnmarshal(data, &member); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := member.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
member.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
member.SetExpiry(cfg.Expires)
}
return &member, nil
}
// AccessListReviews defines an interface for managing Access List reviews.
type AccessListReviews interface {
// ListAccessListReviews will list access list reviews for a particular access list.
ListAccessListReviews(ctx context.Context, accessList string, pageSize int, pageToken string) (reviews []*accesslist.Review, nextToken string, err error)
// ListAccessListReviewsV2 will list access list reviews for a particular access list.
ListAccessListReviewsV2(ctx context.Context, req *accesslistv1.ListAccessListReviewsRequest) (reviews []*accesslist.Review, nextToken string, err error)
// ListAllAccessListReviews will list access list reviews for all unscoped access lists. Only to be used by the cache.
ListAllAccessListReviews(ctx context.Context, pageSize int, pageToken string) (reviews []*accesslist.Review, nextToken string, err error)
// ListAllAccessListReviewsV2 will list access list reviews for all access lists. Only to be used by the cache.
ListAllAccessListReviewsV2(ctx context.Context, req *accesslistv1.ListAllAccessListReviewsRequest) (reviews []*accesslist.Review, nextToken string, err error)
// CreateAccessListReview will create a new review for an access list.
CreateAccessListReview(ctx context.Context, review *accesslist.Review) (updatedReview *accesslist.Review, nextReviewDate time.Time, err error)
// DeleteAccessListReview will delete an access list review from the backend.
DeleteAccessListReview(ctx context.Context, accessListName, reviewName string) error
// DeleteAccessListReviewV2 will delete an access list review from the backend.
DeleteAccessListReviewV2(ctx context.Context, req *accesslistv1.DeleteAccessListReviewRequest) error
}
// MarshalAccessListReview marshals the access list review resource to JSON.
func MarshalAccessListReview(review *accesslist.Review, opts ...MarshalOption) ([]byte, error) {
if err := review.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *review
copy.SetRevision("")
review = ©
}
return utils.FastMarshal(review)
}
// UnmarshalAccessListReview unmarshals the access list review resource from JSON.
func UnmarshalAccessListReview(data []byte, opts ...MarshalOption) (*accesslist.Review, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing access list review data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var review accesslist.Review
if err := utils.FastUnmarshal(data, &review); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := review.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
review.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
review.SetExpiry(cfg.Expires)
}
return &review, nil
}
// CreateAccessListNextKey creates a pagination token based on the requested sort index name
func CreateAccessListNextKey(al *accesslist.AccessList, indexName string) (string, error) {
switch indexName {
case "name":
return AccessListNameIndexKey(al), nil
case "auditNextDate":
return AccessListAuditDateIndexKey(al), nil
case "title":
return AccessListTitleIndexKey(al), nil
default:
return "", trace.BadParameter("unsupported sort %s but expected name, title or auditNextDate", indexName)
}
}
// AccessListNameIndexKey returns the resource name returned from GetName().
func AccessListNameIndexKey(al *accesslist.AccessList) string {
return scopes.MakeResourceCursor(al.GetScope(), al.GetName())
}
// AccessListAuditDateIndexKey returns the DateOnly formatted next audit date
// followed by the resource name for disambiguation.
func AccessListAuditDateIndexKey(al *accesslist.AccessList) string {
if !al.IsReviewable() || al.Spec.Audit.NextAuditDate.IsZero() {
// Use last lexical character to ensure that ACLs without an audit date
// appear at the end when sorted. Otherwise we would compare against
// `0001-01-01 00:00:00` which would sort first, but actually means
// the access list is not eligible for review.
return "z/" + AccessListNameIndexKey(al)
}
return al.Spec.Audit.NextAuditDate.Format(time.DateOnly) + "/" + AccessListNameIndexKey(al)
}
// AccessListTitleIndexKey returns the access list title base32hex encoded
// followed by the resource name for disambiguation.
func AccessListTitleIndexKey(al *accesslist.AccessList) string {
title := cases.Fold().String(al.Spec.Title)
title = base32.HexEncoding.WithPadding(base32.NoPadding).EncodeToString([]byte(title))
return title + "/" + AccessListNameIndexKey(al)
}
// AccessListSearchTermMatcherFunc reports whether an access list matches a search term.
type AccessListSearchTermMatcherFunc func(al *accesslist.AccessList, term string) bool
// MatchAccessList returns true if the access list matches the given filter criteria.
// The function applies filters in sequence: owners, then roles, then search.
// All provided filters must match for the access list to be included.
//
// - If owners filter is provided, the access list must have at least one matching owner
// - If roles filter is provided, the access list must grant at least one matching role
// - If search filter is provided, all search terms must be found across the access list's
// title, name, owner names, description, granted roles, and origin fields, or match a search term matcher
//
// All matching is case-insensitive and supports partial matches.
// Search term matchers are evaluated in order for terms that do not match stored access list fields.
func MatchAccessList(al *accesslist.AccessList, req *accesslistv1.AccessListsFilter, searchTermMatchers ...AccessListSearchTermMatcherFunc) bool {
if req == nil {
return true
}
search := req.GetSearch()
owners := req.GetOwners()
roles := req.GetRoles()
origin := req.GetOrigin()
if search == "" && len(owners) == 0 && len(roles) == 0 && origin == "" {
return true
}
// Step 1: Check owner filter if provided
if len(owners) > 0 {
ownerMatched := slices.ContainsFunc(owners, func(filterOwner string) bool {
return slices.ContainsFunc(al.Spec.Owners, func(alOwner accesslist.Owner) bool {
return strcase.Contains(alOwner.Name, filterOwner)
})
})
if !ownerMatched {
return false
}
}
// Step 2: Check role filter if provided
if len(roles) > 0 {
roleMatched := slices.ContainsFunc(roles, func(filterRole string) bool {
return slices.ContainsFunc(al.Spec.Grants.Roles, func(alRole string) bool {
return strcase.Contains(alRole, filterRole)
})
})
if !roleMatched {
return false
}
}
// Step 3: check the origin
if origin != "" && !strcase.Contains(al.Origin(), origin) {
return false
}
// Step 4: Check search filter if provided
if searchTerms := strings.Fields(search); len(searchTerms) > 0 {
// Check if all search terms are found across the access list fields
// without creating an intermediate slice
for _, term := range searchTerms {
termFound := false
// Check title
if strcase.Contains(al.Spec.Title, term) {
termFound = true
}
// Check name
if !termFound && strcase.Contains(al.GetName(), term) {
termFound = true
}
// Check description
if !termFound && strcase.Contains(al.Spec.Description, term) {
termFound = true
}
// Check origin
if !termFound && strcase.Contains(al.Origin(), term) {
termFound = true
}
// Check owner names
if !termFound {
for _, owner := range al.Spec.Owners {
if strcase.Contains(owner.Name, term) {
termFound = true
break
}
}
}
// Check roles
if !termFound {
for _, role := range al.Spec.Grants.Roles {
if strcase.Contains(role, term) {
termFound = true
break
}
}
}
if !termFound {
termFound = slices.ContainsFunc(searchTermMatchers, func(matcher AccessListSearchTermMatcherFunc) bool {
return matcher(al, term)
})
}
// If this term wasn't found in any field, the search fails
if !termFound {
return false
}
}
}
return true
}
/*
* Teleport
* Copyright (C) 2024 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"
"slices"
"time"
_ "time/tzdata"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/defaults"
accessmonitoringrulesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/accessmonitoringrules/v1"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/accessmonitoring"
"github.com/gravitational/teleport/lib/utils/typical"
)
var (
// accessRequestConditionParser is a parser for the access request condition.
// It is used to validate access monitoring rules before write operations.
accessRequestConditionParser = mustNewAccessRequestConditionParser()
)
// AccessMonitoringRules is the AccessMonitoringRule service
type AccessMonitoringRules interface {
CreateAccessMonitoringRule(ctx context.Context, in *accessmonitoringrulesv1.AccessMonitoringRule) (*accessmonitoringrulesv1.AccessMonitoringRule, error)
UpdateAccessMonitoringRule(ctx context.Context, in *accessmonitoringrulesv1.AccessMonitoringRule) (*accessmonitoringrulesv1.AccessMonitoringRule, error)
UpsertAccessMonitoringRule(ctx context.Context, in *accessmonitoringrulesv1.AccessMonitoringRule) (*accessmonitoringrulesv1.AccessMonitoringRule, error)
GetAccessMonitoringRule(ctx context.Context, name string) (*accessmonitoringrulesv1.AccessMonitoringRule, error)
DeleteAccessMonitoringRule(ctx context.Context, name string) error
DeleteAllAccessMonitoringRules(ctx context.Context) error
ListAccessMonitoringRules(ctx context.Context, limit int, startKey string) ([]*accessmonitoringrulesv1.AccessMonitoringRule, string, error)
ListAccessMonitoringRulesWithFilter(ctx context.Context, req *accessmonitoringrulesv1.ListAccessMonitoringRulesWithFilterRequest) ([]*accessmonitoringrulesv1.AccessMonitoringRule, string, error)
}
// NewAccessMonitoringRuleWithLabels creates a new AccessMonitoringRule with the given spec and labels.
func NewAccessMonitoringRuleWithLabels(name string, labels map[string]string, spec *accessmonitoringrulesv1.AccessMonitoringRuleSpec) (*accessmonitoringrulesv1.AccessMonitoringRule, error) {
amr := accessmonitoringrulesv1.AccessMonitoringRule_builder{
Kind: types.KindAccessMonitoringRule,
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: name,
Namespace: defaults.Namespace,
Labels: labels,
}.Build(),
Spec: spec,
}.Build()
err := ValidateAccessMonitoringRule(amr)
if err != nil {
return nil, trace.Wrap(err)
}
return amr, nil
}
// ValidateAccessMonitoringRule checks that the provided access monitoring rule is valid.
func ValidateAccessMonitoringRule(accessMonitoringRule *accessmonitoringrulesv1.AccessMonitoringRule) error {
if accessMonitoringRule.GetKind() != types.KindAccessMonitoringRule {
return trace.BadParameter("invalid kind for access monitoring rule: %q", accessMonitoringRule.GetKind())
}
if !accessMonitoringRule.HasMetadata() {
return trace.BadParameter("accessMonitoringRule metadata is missing")
}
if accessMonitoringRule.GetVersion() != types.V1 {
return trace.BadParameter("accessMonitoringRule version %q is not supported", accessMonitoringRule.GetVersion())
}
if !accessMonitoringRule.HasSpec() {
return trace.BadParameter("accessMonitoringRule spec is missing")
}
if len(accessMonitoringRule.GetSpec().GetSubjects()) == 0 {
return trace.BadParameter("accessMonitoringRule subject is missing")
}
if accessMonitoringRule.GetSpec().GetCondition() == "" {
return trace.BadParameter("accessMonitoringRule condition is missing")
}
if err := validateSchedules(accessMonitoringRule.GetSpec().GetSchedules()); err != nil {
return trace.Wrap(err, "validating spec.schedules")
}
if accessMonitoringRule.GetSpec().HasNotification() && accessMonitoringRule.GetSpec().GetNotification().GetName() == "" {
return trace.BadParameter("accessMonitoringRule notification plugin name is missing")
}
if automaticReview := accessMonitoringRule.GetSpec().GetAutomaticReview(); automaticReview != nil {
if automaticReview.GetIntegration() == "" {
return trace.BadParameter("accessMonitoringRule automatic_review integration is missing")
}
switch automaticReview.GetDecision() {
case types.RequestState_APPROVED.String(), types.RequestState_DENIED.String():
case "":
return trace.BadParameter("accessMonitoringRule automatic_review decision is missing")
default:
return trace.BadParameter("accessMonitoringRule automatic_review decision %q is not supported", automaticReview.GetDecision())
}
// The automatic review reason becomes the access request's resolve reason
// when the review resolves the request, so it must satisfy the access
// request reason limit
if len(automaticReview.GetReason()) > maxAccessRequestReasonSize {
return trace.BadParameter("accessMonitoringRule automatic_review reason is too long, max %v bytes", maxAccessRequestReasonSize)
}
}
if slices.Contains(accessMonitoringRule.GetSpec().GetSubjects(), types.KindAccessRequest) {
_, err := accessRequestConditionParser.Parse(accessMonitoringRule.GetSpec().GetCondition())
if err != nil {
return trace.BadParameter("accessMonitoringRule condition is invalid: %s", err.Error())
}
desiredState := accessMonitoringRule.GetSpec().GetDesiredState()
switch desiredState {
case "", types.AccessMonitoringRuleStateReviewed:
default:
return trace.BadParameter("accessMonitoringRule desired_state %q is not supported", desiredState)
}
if accessMonitoringRule.GetSpec().GetNotification() != nil {
return nil
}
if accessMonitoringRule.GetSpec().GetAutomaticReview() != nil {
return nil
}
return trace.BadParameter("one of notification or automatic_review must be configured if subjects contain %q",
types.KindAccessRequest)
}
return nil
}
func validateSchedules(schedules map[string]*accessmonitoringrulesv1.Schedule) error {
for name, schedule := range schedules {
if schedule.GetTime() == nil {
return trace.BadParameter("spec.schedules[%s].time is required", name)
}
if err := validateTimeSchedule(schedule.GetTime()); err != nil {
return trace.Wrap(err, "validating spec.schedules[%s].time", name)
}
}
return nil
}
func validateTimeSchedule(schedule *accessmonitoringrulesv1.TimeSchedule) error {
if _, err := time.LoadLocation(schedule.GetTimezone()); err != nil {
return trace.Wrap(err, "invalid timezone: refer to the IANA Time Zone Database for valid options")
}
if len(schedule.GetShifts()) == 0 {
return trace.BadParameter("at least one shift is required")
}
for _, shift := range schedule.GetShifts() {
if err := validateShift(shift); err != nil {
return trace.Wrap(err, "shift is invalid")
}
}
return nil
}
func validateShift(shift *accessmonitoringrulesv1.TimeSchedule_Shift) error {
if _, ok := types.ParseWeekday(shift.GetWeekday()); !ok {
return trace.BadParameter("failed to parse weekday: %v", shift.GetWeekday())
}
start, err := accessmonitoring.ClockTime(time.Time{}, shift.GetStart())
if err != nil {
return trace.Wrap(err, "invalid start time")
}
end, err := accessmonitoring.ClockTime(time.Time{}, shift.GetEnd())
if err != nil {
return trace.Wrap(err, "invalid end time")
}
if !start.Before(end) {
return trace.BadParameter("start time must be before end time")
}
return nil
}
// MarshalAccessMonitoringRule marshals AccessMonitoringRule resource to JSON.
func MarshalAccessMonitoringRule(accessMonitoringRule *accessmonitoringrulesv1.AccessMonitoringRule, opts ...MarshalOption) ([]byte, error) {
return FastMarshalProtoResourceDeprecated(accessMonitoringRule, opts...)
}
// UnmarshalAccessMonitoringRule unmarshals the AccessMonitoringRule resource.
func UnmarshalAccessMonitoringRule(data []byte, opts ...MarshalOption) (*accessmonitoringrulesv1.AccessMonitoringRule, error) {
return FastUnmarshalProtoResourceDeprecated[*accessmonitoringrulesv1.AccessMonitoringRule](data, opts...)
}
// MatchAccessMonitoringRule returns true if the provided rule matches the provided match fields.
// The match fields are optional. If a match field is not provided, then the
// rule matches any value for that field.
func MatchAccessMonitoringRule(rule *accessmonitoringrulesv1.AccessMonitoringRule, subjects []string, notificationIntegration, automaticReviewIntegration string) bool {
if notificationIntegration != "" {
if rule.GetSpec().GetNotification().GetName() != notificationIntegration {
return false
}
}
if automaticReviewIntegration != "" {
if rule.GetSpec().GetAutomaticReview().GetIntegration() != automaticReviewIntegration {
return false
}
}
for _, subject := range subjects {
if ok := slices.ContainsFunc(rule.GetSpec().GetSubjects(), func(s string) bool {
return s == subject
}); !ok {
return false
}
}
return true
}
func mustNewAccessRequestConditionParser() *typical.Parser[accessmonitoring.AccessRequestExpressionEnv, any] {
parser, err := accessmonitoring.NewAccessRequestConditionParser()
if err != nil {
panic(err)
}
return parser
}
/*
* 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"
"fmt"
"log/slog"
"maps"
"slices"
"sort"
"strings"
"time"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"k8s.io/apimachinery/pkg/runtime/schema"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/accessrequest"
"github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/parse"
"github.com/gravitational/teleport/lib/utils/set"
"github.com/gravitational/teleport/lib/utils/typical"
)
const (
maxAccessRequestReasonSize = 4096
maxResourcesPerRequest = 300
maxResourcesLength = 2048
// A day is sometimes 23 hours, sometimes 25 hours, usually 24 hours.
day = 24 * time.Hour
// MaxAccessDuration is the maximum duration that an access request can be
// granted for.
MaxAccessDuration = 14 * day
// requestTTL is the TTL for an access request, i.e. the amount of time that
// the access request can be reviewed. Defaults to 1 week.
requestTTL = 7 * day
// InvalidKubernetesKindAccessRequest is used in part of error messages related to
// `request.kubernetes_resources` config. It's also used to determine if a returned error
// contains this string (in tests and tsh) to customize error messages shown to user.
InvalidKubernetesKindAccessRequest = `your Teleport role's "request.kubernetes_resources" field`
// CannotRequestRole is used in error messages when a user attempts to request
// a role they are not allowed to use. It is also checked in returned errors
// (in tests and tsh) to customize the error message shown to the user.
CannotRequestRole = "can not request role"
)
// ValidateAccessRequest validates the AccessRequest and sets default values
func ValidateAccessRequest(ar types.AccessRequest) error {
return validateAccessRequest(ar, false)
}
// validateAccessRequest implements [ValidateAccessRequest]. With
// allowUnenforceable set, unenforceable constraints (see
// [types.ResourceConstraints.Unenforceable]) pass validation, keeping
// requests written by newer Auths readable.
func validateAccessRequest(ar types.AccessRequest, allowUnenforceable bool) error {
if err := CheckAndSetDefaults(ar); err != nil {
return trace.Wrap(err)
}
_, err := uuid.Parse(ar.GetName())
if err != nil {
return trace.BadParameter("invalid access request ID %q", ar.GetName())
}
if len(ar.GetRequestReason()) > maxAccessRequestReasonSize {
return trace.BadParameter("access request reason is too long, max %v bytes", maxAccessRequestReasonSize)
}
if len(ar.GetResolveReason()) > maxAccessRequestReasonSize {
return trace.BadParameter("access request resolve reason is too long, max %v bytes", maxAccessRequestReasonSize)
}
if l := len(ar.GetAllRequestedResourceIDs()); l > maxResourcesPerRequest {
return trace.BadParameter("access request contains too many resources (%v), max %v", l, maxResourcesPerRequest)
}
for _, r := range ar.GetRequestedResourceAccessIDs() {
rc := r.GetConstraints()
if rc == nil {
continue
}
if allowUnenforceable && rc.Unenforceable() {
// Skip all validation; these may carry a newer version.
continue
}
if err := rc.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
kind := r.GetResourceID().Kind
switch c := rc.Details.(type) {
case *types.ResourceConstraints_AwsConsole:
if kind != types.KindApp {
return trace.BadParameter("aws_console constraints are not valid for resource kind %q", kind)
}
case *types.ResourceConstraints_Ssh:
if kind != types.KindNode {
return trace.BadParameter("ssh constraints are not valid for resource kind %q", kind)
}
default:
return trace.BadParameter("unsupported constraint type %T for resource kind %q", c, kind)
}
}
return nil
}
// ClusterGetter provides access to the local cluster
type ClusterGetter interface {
ClusterNameGetter
// GetRemoteCluster returns a remote cluster by name
GetRemoteCluster(ctx context.Context, clusterName string) (types.RemoteCluster, error)
}
// ValidateAccessRequestClusterNames checks that the clusters in the access request exist
func ValidateAccessRequestClusterNames(cg ClusterGetter, ar types.AccessRequest) error {
ctx := context.TODO()
localClusterName, err := cg.GetClusterName(ctx)
if err != nil {
return trace.Wrap(err)
}
var invalidClusters []string
for _, resourceAccessID := range ar.GetAllRequestedResourceIDs() {
resourceID := resourceAccessID.GetResourceID()
if resourceID.ClusterName == "" {
continue
}
if resourceID.ClusterName == localClusterName.GetClusterName() {
continue
}
_, err := cg.GetRemoteCluster(ctx, resourceID.ClusterName)
if err != nil && !trace.IsNotFound(err) {
return trace.Wrap(err, "failed to fetch remote cluster %q", resourceID.ClusterName)
}
if trace.IsNotFound(err) {
invalidClusters = append(invalidClusters, resourceID.ClusterName)
}
}
if len(invalidClusters) > 0 {
return trace.NotFound("access request contains invalid or unknown cluster names: %v",
strings.Join(apiutils.Deduplicate(invalidClusters), ", "))
}
return nil
}
// NewAccessRequest assembles an AccessRequest resource.
func NewAccessRequest(user string, roles ...string) (types.AccessRequest, error) {
return NewAccessRequestWithResources(user, roles, []types.ResourceAccessID{})
}
// NewAccessRequestWithResources assembles an AccessRequest resource with
// requested resources.
func NewAccessRequestWithResources(user string, roles []string, resourceIDs []types.ResourceAccessID) (types.AccessRequest, error) {
req, err := types.NewAccessRequestWithResources(uuid.New().String(), user, roles, resourceIDs)
if err != nil {
return nil, trace.Wrap(err)
}
if err := ValidateAccessRequest(req); err != nil {
return nil, trace.Wrap(err)
}
return req, nil
}
// AccessRequestGetter defines the interface for fetching access request resources.
type AccessRequestGetter interface {
// GetAccessRequests gets all currently active access requests.
GetAccessRequests(ctx context.Context, filter types.AccessRequestFilter) ([]types.AccessRequest, error)
// ListAccessRequests is an access request getter with pagination and sorting options.
ListAccessRequests(ctx context.Context, req *proto.ListAccessRequestsRequest) (*proto.ListAccessRequestsResponse, error)
}
// DynamicAccessCore is the core functionality common to all DynamicAccess implementations.
type DynamicAccessCore interface {
AccessRequestGetter
// CreateAccessRequestV2 stores a new access request.
CreateAccessRequestV2(ctx context.Context, req types.AccessRequest) (types.AccessRequest, error)
// DeleteAccessRequest deletes an access request.
DeleteAccessRequest(ctx context.Context, reqID string) error
}
// DynamicAccess is a service which manages dynamic RBAC. Specifically, this is the
// dynamic access interface implemented by remote clients.
type DynamicAccess interface {
DynamicAccessCore
// SetAccessRequestState updates the state of an existing access request.
SetAccessRequestState(ctx context.Context, params types.AccessRequestUpdate) error
// SubmitAccessReview applies a review to a request and returns the post-application state.
SubmitAccessReview(ctx context.Context, params types.AccessReviewSubmission) (types.AccessRequest, error)
// GetAccessRequestAllowedPromotions returns suggested access lists for the given access request.
GetAccessRequestAllowedPromotions(ctx context.Context, req types.AccessRequest) (*types.AccessRequestAllowedPromotions, error)
}
// DynamicAccessOracle is a service capable of answering questions related
// to the dynamic access API. Necessary because some information (e.g. the
// list of roles a user is allowed to request) can not be calculated by
// actors with limited privileges.
type DynamicAccessOracle interface {
GetAccessCapabilities(ctx context.Context, req types.AccessCapabilitiesRequest) (*types.AccessCapabilities, error)
GetRemoteAccessCapabilities(ctx context.Context, req types.RemoteAccessCapabilitiesRequest) (*types.RemoteAccessCapabilities, error)
}
func shouldFilterRequestableRolesByResource(a RequestValidatorGetter, req types.AccessCapabilitiesRequest) (bool, error) {
if !req.FilterRequestableRolesByResource {
return false, nil
}
currentCluster, err := a.GetClusterName(context.TODO())
if err != nil {
return false, trace.Wrap(err)
}
for _, resourceAccessID := range types.CombineAsResourceAccessIDs(req.ResourceIDs, req.ResourceAccessIds) {
if resourceAccessID.GetResourceID().ClusterName != currentCluster.GetClusterName() {
// Requested resource is from another cluster, so we can't know
// all of the roles which would grant access to it.
return false, nil
}
}
return true, nil
}
// CalculateAccessCapabilities aggregates the requested capabilities using the supplied getter
// to load relevant resources.
func CalculateAccessCapabilities(ctx context.Context, clock clockwork.Clock, clt RequestValidatorGetter, identity tlsca.Identity, req types.AccessCapabilitiesRequest) (*types.AccessCapabilities, error) {
shouldFilter, err := shouldFilterRequestableRolesByResource(clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
if !shouldFilter && req.FilterRequestableRolesByResource {
req.ResourceIDs = nil
req.ResourceAccessIds = nil
}
var caps types.AccessCapabilities
// all capabilities require use of a request validator. calculating suggested reviewers
// requires that the validator be configured for variable expansion.
v, err := NewRequestValidator(ctx, clock, clt, req.User, WithExpandVars(req.SuggestedReviewers))
if err != nil {
return nil, trace.Wrap(err)
}
resourceAccessIDs := types.CombineAsResourceAccessIDs(req.ResourceIDs, req.ResourceAccessIds)
if len(resourceAccessIDs) != 0 && !req.FilterRequestableRolesByResource {
caps.ApplicableRolesForResources, err = v.applicableSearchAsRoles(ctx, resourceAccessIDs, req.Login)
if err != nil {
return nil, trace.Wrap(err)
}
}
if req.RequestableRoles {
var requestableResourceAccessIDs []types.ResourceAccessID
if req.FilterRequestableRolesByResource {
requestableResourceAccessIDs = resourceAccessIDs
}
caps.RequestableRoles, err = v.getRequestableRoles(ctx, identity, requestableResourceAccessIDs, req.Login)
if err != nil {
return nil, trace.Wrap(err)
}
}
if req.SuggestedReviewers {
caps.SuggestedReviewers = v.suggestedReviewers
}
caps.RequireReason, err = v.calcRequireReasonCap(ctx, req, caps)
if err != nil {
return nil, trace.Wrap(err)
}
if len(v.reasonPrompts) > 0 {
caps.RequestPrompt = v.reasonPrompts[0]
}
caps.AutoRequest = v.autoRequestOnLogin
return &caps, nil
}
// PruneMappedSearchAsRoles calculates the roles required to access the given
// resources based on the supplied set of `search_as` roles.
func PruneMappedSearchAsRoles(ctx context.Context, clock clockwork.Clock, getter RequestValidatorGetter, mappedUser UserState, mappedSearchAsRoles []string, resourceAccessIDs []types.ResourceAccessID, loginHint string) ([]string, error) {
clusterNameResource, err := getter.GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
localClusterName := clusterNameResource.GetClusterName()
hasRemoteResources := slices.ContainsFunc(resourceAccessIDs,
func(rID types.ResourceAccessID) bool {
return rID.GetResourceID().ClusterName != localClusterName
})
if hasRemoteResources {
return nil, trace.BadParameter("request must only contain resources in local cluster")
}
rv, err := NewRequestValidatorForUser(ctx, clock, getter, mappedUser)
if err != nil {
return nil, trace.Wrap(err)
}
roles, err := rv.pruneResourceRequestRoles(ctx, resourceAccessIDs, loginHint, mappedSearchAsRoles)
if err != nil {
return nil, trace.Wrap(err)
}
return roles, nil
}
func (v *RequestValidator) calcRequireReasonCap(ctx context.Context, req types.AccessCapabilitiesRequest, caps types.AccessCapabilities) (requireReason bool, err error) {
var roles []string
if req.RequestableRoles {
roles = caps.RequestableRoles
} else {
roles = caps.ApplicableRolesForResources
}
requireReason, _, err = v.isReasonRequired(ctx, roles, nil)
if err != nil {
return false, trace.Wrap(err)
}
return requireReason, nil
}
// allowedSearchAsRoles returns all allowed `allow.request.search_as_roles` for the user that are
// not in the `deny.request.search_as_roles`. It does not filter out any roles that should not be
// allowed based on requests.
func (m *RequestValidator) allowedSearchAsRoles() ([]string, error) {
var rolesToRequest []string
for _, roleName := range m.roles.allowSearch {
if !m.canSearchAsRole(roleName) {
continue
}
rolesToRequest = append(rolesToRequest, roleName)
}
if len(rolesToRequest) == 0 {
return nil, trace.AccessDenied(`Resource Access Requests require usable "search_as_roles", none found for user %q`, m.userState.GetName())
}
return rolesToRequest, nil
}
// applicableSearchAsRoles prunes the search_as_roles and only returns those
// applicable for the given list of resourceIDs.
//
// If loginHint is provided, it will attempt to prune the list to a single role.
func (m *RequestValidator) applicableSearchAsRoles(ctx context.Context, resourceAccessIDs []types.ResourceAccessID, loginHint string) ([]string, error) {
rolesToRequest, err := m.allowedSearchAsRoles()
if err != nil {
return nil, trace.Wrap(err)
}
// Prune the list of roles to request to only those which may be necessary
// to access the requested resources.
rolesToRequest, err = m.pruneResourceRequestRoles(ctx, resourceAccessIDs, loginHint, rolesToRequest)
if err != nil {
return nil, trace.Wrap(err)
}
return rolesToRequest, nil
}
// DynamicAccessExt is an extended dynamic access interface
// used to implement some auth server internals.
type DynamicAccessExt interface {
DynamicAccessCore
// CreateAccessRequest stores a new access request.
CreateAccessRequest(ctx context.Context, req types.AccessRequest) error
// ApplyAccessReview applies a review to a request in the backend and returns the post-application state.
ApplyAccessReview(ctx context.Context, params types.AccessReviewSubmission, checker ReviewPermissionChecker) (types.AccessRequest, error)
// UpsertAccessRequest creates or updates an access request.
UpsertAccessRequest(ctx context.Context, req types.AccessRequest) error
// DeleteAllAccessRequests deletes all existent access requests.
DeleteAllAccessRequests(ctx context.Context) error
// SetAccessRequestState updates the state of an existing access request.
SetAccessRequestState(ctx context.Context, params types.AccessRequestUpdate) (types.AccessRequest, error)
// CreateAccessRequestAllowedPromotions creates a list of allowed access list promotions for the given access request.
CreateAccessRequestAllowedPromotions(ctx context.Context, req types.AccessRequest, accessLists *types.AccessRequestAllowedPromotions) error
// GetAccessRequestAllowedPromotions returns a lists of allowed access list promotions for the given access request.
GetAccessRequestAllowedPromotions(ctx context.Context, req types.AccessRequest) (*types.AccessRequestAllowedPromotions, error)
// ListExpiredAccessRequests lists all access requests that are expired. This is used by
// the expiry service. Access requests expiration handling is done outside the backend
// because we need to emit audit events on the access requests expiry.
ListExpiredAccessRequests(ctx context.Context, limit int, pageToken string) ([]*types.AccessRequestV3, string, error)
}
// reviewParamsContext is a simplified view of an access review
// which represents the incoming review during review threshold
// filter evaluation.
type reviewParamsContext struct {
reason string
annotations map[string][]string
}
// reviewAuthorContext is a simplified view of a user
// resource which represents the author of a review during
// review threshold filter evaluation.
type reviewAuthorContext struct {
roles []string
traits map[string][]string
}
// reviewRequestContext is a simplified view of an access request
// resource which represents the request parameters which are in-scope
// during review threshold filter evaluation.
type reviewRequestContext struct {
roles []string
reason string
systemAnnotations map[string][]string
}
// thresholdFilterContext is the top-level context used to evaluate
// review threshold filters.
type thresholdFilterContext struct {
reviewer reviewAuthorContext
review reviewParamsContext
request reviewRequestContext
}
// reviewPermissionContext is the top-level context used to evaluate
// a user's review permissions. It is functionally identical to the
// thresholdFilterContext except that it does not expose review parameters.
// This is because review permissions are used to determine which requests
// a user is allowed to see, and therefore needs to be calculable prior
// to construction of review parameters.
type reviewPermissionContext struct {
reviewer reviewAuthorContext
request reviewRequestContext
}
// ValidateAccessPredicates checks request & review permission predicates for
// syntax errors. Used to help prevent users from accidentally writing incorrect
// predicates. This function should only be called by the auth server prior to
// storing new/updated roles. Normal role validation deliberately omits these
// checks to allow us to extend the available namespaces without breaking
// backwards compatibility with older nodes/proxies (which never need to evaluate
// these predicates).
func ValidateAccessPredicates(role types.Role) error {
var errs []error
if len(role.GetAccessRequestConditions(types.Deny).Thresholds) != 0 {
// deny blocks never contain thresholds. a threshold which happens to describe a *denial condition* is
// still part of the "allow" block. thresholds are not part of deny blocks because thresholds describe the
// state-transition scenarios supported by a request (including potentially being denied). deny.request blocks match
// requests which are *never* allowable, and therefore will never reach the point of needing to encode thresholds.
errs = append(errs, trace.BadParameter("deny.request cannot contain thresholds, set denial counts in allow.request.thresholds instead"))
}
for i, t := range role.GetAccessRequestConditions(types.Allow).Thresholds {
if t.Filter == "" {
continue
}
if _, err := parseThresholdFilterExpression(t.Filter); err != nil {
errs = append(errs, trace.BadParameter("invalid threshold predicate at allow.request.thresholds[%d]: %q, %v", i, t.Filter, err))
}
}
if w := role.GetAccessReviewConditions(types.Deny).Where; w != "" {
if _, err := parseReviewPermissionExpression(w); err != nil {
errs = append(errs, trace.BadParameter("invalid review predicate at deny.review_requests.where: %q, %v", w, err))
}
}
if w := role.GetAccessReviewConditions(types.Allow).Where; w != "" {
if _, err := parseReviewPermissionExpression(w); err != nil {
errs = append(errs, trace.BadParameter("invalid review predicate at allow.review_requests.where: %q, %v", w, err))
}
}
if maxDuration := role.GetAccessRequestConditions(types.Allow).MaxDuration; maxDuration.Duration() != 0 &&
maxDuration.Duration() > MaxAccessDuration {
errs = append(errs, trace.BadParameter("max access duration must be less than or equal to %v", MaxAccessDuration))
}
return trace.NewAggregate(errs...)
}
// ApplyAccessReview attempts to apply the specified access review to the specified request.
func ApplyAccessReview(req types.AccessRequest, rev types.AccessReview, author UserState) error {
if rev.Author != author.GetName() {
return trace.BadParameter("mismatched review author (expected %q, got %q)", rev.Author, author)
}
// role lists must be deduplicated and sorted
rev.Roles = apiutils.Deduplicate(rev.Roles)
sort.Strings(rev.Roles)
// basic compatibility/sanity checks
if err := checkReviewCompat(req, rev); err != nil {
return trace.Wrap(err)
}
// aggregate the threshold indexes for this review
tids, err := collectReviewThresholdIndexes(req, rev, author)
if err != nil {
return trace.Wrap(err)
}
// set threshold indexes
rev.ThresholdIndexes = tids
// set a review created time if not already set
if rev.Created.IsZero() {
rev.Created = time.Now()
}
// Resolved requests should not be updated.
switch {
case req.GetState().IsApproved():
return trace.AccessDenied("the access request has been already approved")
case req.GetState().IsDenied():
return trace.AccessDenied("the access request has been already denied")
case req.GetState().IsPromoted():
return trace.AccessDenied("the access request has been already promoted")
}
req.SetReviews(append(req.GetReviews(), rev))
if rev.AssumeStartTime != nil {
if err := types.ValidateAssumeStartTime(*rev.AssumeStartTime, req.GetAccessExpiry(), req.GetCreationTime()); err != nil {
return trace.Wrap(err)
}
req.SetAssumeStartTime(*rev.AssumeStartTime)
}
// the request is still pending, so check to see if this
// review introduces a state-transition.
res, err := calculateReviewBasedResolution(req)
if err != nil || res == nil {
return trace.Wrap(err)
}
// state-transition was triggered. update the appropriate fields.
if err := req.SetState(res.state); err != nil {
return trace.Wrap(err)
}
req.SetResolveReason(res.reason)
if req.GetPromotedAccessListName() == "" {
// Set the title only if it's not set yet. This is to prevent
// overwriting the title by another promotion review.
req.SetPromotedAccessListName(rev.GetAccessListName())
req.SetPromotedAccessListTitle(rev.GetAccessListTitle())
}
req.SetExpiry(req.GetAccessExpiry())
return nil
}
// checkReviewCompat performs basic checks to ensure that the specified review can be
// applied to the specified request (part of review application logic).
func checkReviewCompat(req types.AccessRequest, rev types.AccessReview) error {
// The Proposal cannot be yet resolved.
if !rev.ProposedState.IsResolved() {
// Skip the promoted state in the error message. It's not a state that most people
// should be concerned with.
return trace.BadParameter("invalid state proposal: %s (expected approval/denial)", rev.ProposedState)
}
// the default threshold should exist. if it does not, the request either is not fully
// initialized (i.e., variable expansion has not been run yet), or the request was inserted into
// the backend by a teleport instance which does not support the review feature.
if len(req.GetThresholds()) == 0 {
return trace.BadParameter("request is uninitialized or does not support reviews")
}
// A review submitted by an identity (eg. plugin), for another user, cannot be applied to the submitter's
// own request.
if rev.SubmittedBy == req.GetUser() {
return trace.AccessDenied("review submitter %q cannot apply a review on their own request", rev.SubmittedBy)
}
// user must not have previously reviewed this request
for _, existingReview := range req.GetReviews() {
if existingReview.Author == rev.Author {
return trace.AlreadyExists("user %q has already reviewed this request", rev.Author)
}
}
rtm := req.GetRoleThresholdMapping()
// TODO(fspmarshall): Remove this restriction once role overrides
// in reviews are fully supported.
if len(rev.Roles) != 0 && len(rev.Roles) != len(rtm) {
return trace.NotImplemented("role subselection is not yet supported in reviews, try omitting role list")
}
// TODO(fspmarhsall): Remove this restriction once annotations
// in reviews are fully supported.
if len(rev.Annotations) != 0 {
return trace.NotImplemented("annotations are not yet supported in reviews, try omitting annotations field")
}
// verify that all roles are present within the request
for _, role := range rev.Roles {
if _, ok := rtm[role]; !ok {
return trace.BadParameter("role %q is not a member of this request", role)
}
}
return nil
}
// collectReviewThresholdIndexes aggregates the indexes of all thresholds whose filters match
// the supplied review (part of review application logic).
func collectReviewThresholdIndexes(req types.AccessRequest, rev types.AccessReview, author UserState) ([]uint32, error) {
var tids []uint32
ctx := newThresholdFilterContext(req, rev, author)
for i, t := range req.GetThresholds() {
match, err := accessReviewThresholdMatchesFilter(t, ctx)
if err != nil {
return nil, trace.Wrap(err)
}
if !match {
continue
}
tid := uint32(i)
if int(tid) != i {
// sanity-check. we disallow extremely large threshold lists elsewhere, but it's always
// best to double-check these things.
return nil, trace.Errorf("threshold index %d out of supported range (this is a bug)", i)
}
tids = append(tids, tid)
}
return tids, nil
}
// accessReviewThresholdMatchesFilter returns true if Filter rule matches
// Empty Filter block always matches
func accessReviewThresholdMatchesFilter(t types.AccessReviewThreshold, ctx thresholdFilterContext) (bool, error) {
if t.Filter == "" {
return true, nil
}
expr, err := parseThresholdFilterExpression(t.Filter)
if err != nil {
return false, trace.Wrap(err)
}
return expr.Evaluate(ctx)
}
// newThresholdFilterContext creates a custom parser context which exposes a simplified view of the review author
// and the request for evaluation of review threshold filters.
func newThresholdFilterContext(req types.AccessRequest, rev types.AccessReview, author UserState) thresholdFilterContext {
return thresholdFilterContext{
reviewer: reviewAuthorContext{
roles: author.GetRoles(),
traits: author.GetTraits(),
},
review: reviewParamsContext{
reason: rev.Reason,
annotations: rev.Annotations,
},
request: reviewRequestContext{
roles: req.GetOriginalRoles(),
reason: req.GetRequestReason(),
systemAnnotations: req.GetSystemAnnotations(),
},
}
}
// requestResolution describes a request state-transition from
// PENDING to some other state.
type requestResolution struct {
state types.RequestState
reason string
}
// calculateReviewBasedResolution calculates the request resolution based upon
// a request's reviews. Returns (nil,nil) in the event no resolution has been reached.
func calculateReviewBasedResolution(req types.AccessRequest) (*requestResolution, error) {
// thresholds and reviews must be populated before state-transitions are possible
thresholds, reviews := req.GetThresholds(), req.GetReviews()
if len(thresholds) == 0 || len(reviews) == 0 {
return nil, nil
}
// approved keeps track of roles that have hit at least one
// of their approval thresholds.
approved := make(map[string]struct{})
// denied keeps track of whether we've seen *any* role get denied
// (which role does not currently matter since we short-circuit on the
// first denial to be triggered).
denied := false
// counts keeps track of the approval and denial counts for all thresholds.
counts := make([]struct{ approval, denial uint32 }, len(thresholds))
// lastReview stores the most recently processed review. Since processing halts
// once we hit our first approval/denial condition, this review represents the
// triggering review for the approval/denial state-transition.
var lastReview types.AccessReview
// Iterate through all reviews and aggregate them against `counts`.
ProcessReviews:
for _, rev := range reviews {
lastReview = rev
for _, tid := range rev.ThresholdIndexes {
idx := int(tid)
if len(thresholds) <= idx {
return nil, trace.Errorf("threshold index '%d' out of range (this is a bug)", idx)
}
switch {
case rev.ProposedState.IsApproved():
counts[idx].approval++
case rev.ProposedState.IsDenied():
counts[idx].denial++
case rev.ProposedState.IsPromoted():
// Promote skips the threshold check.
break ProcessReviews
default:
return nil, trace.BadParameter("cannot calculate state-transition, unexpected proposal: %s", rev.ProposedState)
}
}
// If we hit any denial thresholds, short-circuit immediately
for i, t := range thresholds {
if counts[i].denial >= t.Deny && t.Deny != 0 {
denied = true
break ProcessReviews
}
}
// check for roles that can be transitioned to an approved state
CheckRoleApprovals:
for role, thresholdSets := range req.GetRoleThresholdMapping() {
if _, ok := approved[role]; ok {
// the role was marked approved during a previous iteration
continue CheckRoleApprovals
}
// iterate through all threshold sets. All sets must have at least
// one threshold which has hit its approval count in order for the
// role to be considered approved.
CheckThresholdSets:
for _, tset := range thresholdSets.Sets {
for _, tid := range tset.Indexes {
idx := int(tid)
if len(thresholds) <= idx {
return nil, trace.Errorf("threshold index out of range %s/%d (this is a bug)", role, tid)
}
t := thresholds[idx]
if counts[idx].approval >= t.Approve && t.Approve != 0 {
// this set contains a threshold which has met its approval condition.
// skip to the next set.
continue CheckThresholdSets
}
}
// no thresholds met for this set. there may be additional roles/thresholds
// that did meet their requirements this iteration, but there is no point in
// processing them unless this set has also hit its requirements. we therefore
// move immediately to processing the next review.
continue ProcessReviews
}
// since we skip to the next review as soon as we see a set which has not hit any of its
// approval scenarios, we know that if we get to this point the role must be approved.
approved[role] = struct{}{}
}
// If we got here, then we iterated across all roles in the rtm without hitting any that
// had not met their approval scenario. The request has hit an approved state and further
// reviews will not be processed.
break ProcessReviews
}
switch {
case lastReview.ProposedState.IsApproved():
if len(approved) != len(req.GetRoleThresholdMapping()) {
// processing halted on approval, but not all roles have
// hit their approval thresholds; no state-transition.
return nil, nil
}
case lastReview.ProposedState.IsDenied():
if !denied {
// processing halted on denial, but no roles have hit
// their denial thresholds; no state-transition.
return nil, nil
}
case lastReview.ProposedState.IsPromoted():
// Let the state change. Promoted won't grant any access, meaning it is roughly equivalent to denial.
// But we want to be able to distinguish between promoted and denied in audit logs/UI.
default:
return nil, trace.BadParameter("cannot calculate state-transition, unexpected proposal: %s", lastReview.ProposedState)
}
// processing halted on valid state-transition; return resolution
// based on last review
return &requestResolution{
state: lastReview.ProposedState,
reason: lastReview.Reason,
}, nil
}
// GetAccessRequest is a helper function assists with loading a specific request by ID.
func GetAccessRequest(ctx context.Context, acc DynamicAccessCore, reqID string) (types.AccessRequest, error) {
reqs, err := acc.GetAccessRequests(ctx, types.AccessRequestFilter{
ID: reqID,
})
if err != nil {
return nil, trace.Wrap(err)
}
if len(reqs) < 1 {
return nil, trace.NotFound("no access request matching %q", reqID)
}
return reqs[0], nil
}
// GetTraitMappings gets the AccessRequestConditions' claims as a TraitMappingsSet
func GetTraitMappings(cms []types.ClaimMapping) types.TraitMappingSet {
tm := make([]types.TraitMapping, 0, len(cms))
for _, mapping := range cms {
tm = append(tm, types.TraitMapping{
Trait: mapping.Claim,
Value: mapping.Value,
Roles: mapping.Roles,
})
}
return types.TraitMappingSet(tm)
}
// RequestValidatorGetter is the interface required by the request validation
// functions used to get the necessary resources.
type RequestValidatorGetter interface {
UserLoginStatesGetter
UserGetter
RoleGetter
client.ListResourcesClient
GetRoles(ctx context.Context) ([]types.Role, error)
GetClusterName(ctx context.Context) (types.ClusterName, error)
}
// AppendRoleMatchers constructs all role matchers for a given
// AccessRequestConditions instance and appends them to the
// supplied matcher slice.
func AppendRoleMatchers(matchers []parse.Matcher, roles []string, cms []types.ClaimMapping, traits map[string][]string) ([]parse.Matcher, error) {
// build matchers for the role list
for _, r := range roles {
m, err := parse.NewMatcher(r)
if err != nil {
return nil, trace.Wrap(err)
}
matchers = append(matchers, m)
}
// build matchers for all role mappings
ms, err := TraitsToRoleMatchers(GetTraitMappings(cms), traits)
if err != nil {
return nil, trace.Wrap(err)
}
return append(matchers, ms...), nil
}
// ReviewPermissionChecker is a helper for validating whether a user
// is allowed to review specific access requests.
type ReviewPermissionChecker struct {
UserState UserState
Roles struct {
// allow/deny mappings sort role matches into lists based on their
// constraining predicate (where) expression.
AllowReview, DenyReview map[string][]parse.Matcher
}
}
// HasAllowDirectives checks if any allow directives exist. A user with
// no allow directives will never be able to review any requests.
func (c *ReviewPermissionChecker) HasAllowDirectives() bool {
for _, allowMatchers := range c.Roles.AllowReview {
if len(allowMatchers) > 0 {
return true
}
}
return false
}
// CanReviewRequest checks if the user is allowed to review the specified request.
// Note that the ability to review a request does not necessarily imply that any specific
// approval/denial thresholds will actually match the user's review. Matching one or more
// thresholds is not a pre-requisite for review submission.
func (c *ReviewPermissionChecker) CanReviewRequest(req types.AccessRequest) (bool, error) {
// TODO(fspmarshall): Refactor this to improve readability when
// adding role subselection support.
// user cannot review their own request
if c.UserState.GetName() == req.GetUser() {
return false, nil
}
// method allocates a new array if an override has already been
// called, so get the role list once in advance.
requestedRoles := req.GetOriginalRoles()
rpc := reviewPermissionContext{
reviewer: reviewAuthorContext{
roles: c.UserState.GetRoles(),
traits: c.UserState.GetTraits(),
},
request: reviewRequestContext{
roles: requestedRoles,
reason: req.GetRequestReason(),
systemAnnotations: req.GetSystemAnnotations(),
},
}
// check all denial rules first.
for expr, denyMatchers := range c.Roles.DenyReview {
// if predicate is non-empty, it must match
if expr != "" {
parsed, err := parseReviewPermissionExpression(expr)
if err != nil {
return false, trace.Wrap(err)
}
match, err := parsed.Evaluate(rpc)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
continue
}
}
for _, role := range requestedRoles {
for _, deny := range denyMatchers {
if deny.Match(role) {
// short-circuit on first denial
return false, nil
}
}
}
}
// needsAllow tracks the list of roles which still need to match an allow directive
// in order for the request to be reviewable. we need to perform a deep copy here
// since we perform a filter-in-place when we find a matching allow directive.
needsAllow := make([]string, len(requestedRoles))
copy(needsAllow, requestedRoles)
Outer:
for expr, allowMatchers := range c.Roles.AllowReview {
// if predicate is non-empty, it must match.
if expr != "" {
parsed, err := parseReviewPermissionExpression(expr)
if err != nil {
return false, trace.Wrap(err)
}
match, err := parsed.Evaluate(rpc)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
continue Outer
}
}
// unmatched collects unmatched roles for our filter-in-place operation.
unmatched := needsAllow[:0]
MatchRoles:
for _, role := range needsAllow {
for _, allow := range allowMatchers {
if allow.Match(role) {
// role matched this allow directive, and will be filtered out
continue MatchRoles
}
}
// still unmatched, this role will continue to be part of
// the needsAllow list next iteration.
unmatched = append(unmatched, role)
}
// finalize our filter-in-place
needsAllow = unmatched
if len(needsAllow) == 0 {
// all roles have matched an allow directive, no further
// processing is required.
break Outer
}
}
return len(needsAllow) == 0, nil
}
type userStateRoleOverride struct {
UserState
Roles []string
}
func (u userStateRoleOverride) GetRoles() []string {
return u.Roles
}
// NewReviewPermissionChecker creates a review permission checker for the Teleport user given
// by the username and identity. The identity is used for bot users that must retain minimal
// permissions granted by the bot identity's roles.
// The caller of this function should verify that the username (review author) and identity
// refer to the same user, or otherwise pass in a nil identity.
func NewReviewPermissionChecker(
ctx context.Context,
getter RequestValidatorGetter,
username string,
identity *tlsca.Identity,
) (ReviewPermissionChecker, error) {
uls, err := GetUserOrLoginState(ctx, getter, username)
if err != nil {
return ReviewPermissionChecker{}, trace.Wrap(err)
}
// By default, the users freshly fetched roles are used rather than the
// roles on the x509 identity. This prevents recursive access request
// review.
//
// For bots, however, the roles on the identity must be used. This is
// because the certs output by a bot always use role impersonation and the
// role directly assigned to a bot has minimal permissions.
if uls.IsBot() {
if identity == nil {
// Handle an edge case where SubmitAccessReview is being invoked
// in-memory but as a bot user.
//
// There should not be a scenario where a different identity (eg. plugin) submits
// a review for a bot user, and this check enforces it.
// Identities submitting for other users should only be able to create
// permission checkers for human users, if they are granted `submit_for_users` permissions.
return ReviewPermissionChecker{}, trace.BadParameter(
"bot user provided but identity parameter is nil",
)
}
if identity.Username != username {
// It should not be possible for these to be different as a
// guard in AuthorizeAccessReviewRequest prevents submitting a
// request as another user unless you have the admin role. This
// safeguard protects against that regressing and creating an
// inconsistent state.
return ReviewPermissionChecker{}, trace.BadParameter(
"bot identity username and review author mismatch",
)
}
if len(identity.ActiveRequests) > 0 {
// It should not be possible for a bot's output certificates to
// have active requests - but this additional check safeguards us
// against a regression elsewhere and prevents recursive access
// requests occurring.
return ReviewPermissionChecker{}, trace.BadParameter(
"bot should not have active requests",
)
}
// Override list of roles to roles currently present on the x509 ident.
uls = userStateRoleOverride{
UserState: uls,
Roles: identity.Groups,
}
}
c := ReviewPermissionChecker{
UserState: uls,
}
c.Roles.AllowReview = make(map[string][]parse.Matcher)
c.Roles.DenyReview = make(map[string][]parse.Matcher)
// load all statically assigned roles for the user and
// use them to build our checker state.
for _, roleName := range c.UserState.GetRoles() {
role, err := getter.GetRole(ctx, roleName)
if err != nil {
return ReviewPermissionChecker{}, trace.Wrap(err)
}
if err := c.push(role); err != nil {
return ReviewPermissionChecker{}, trace.Wrap(err)
}
}
return c, nil
}
func (c *ReviewPermissionChecker) push(role types.Role) error {
allow, deny := role.GetAccessReviewConditions(types.Allow), role.GetAccessReviewConditions(types.Deny)
var err error
c.Roles.DenyReview[deny.Where], err = AppendRoleMatchers(c.Roles.DenyReview[deny.Where], deny.Roles, deny.ClaimsToRoles, c.UserState.GetTraits())
if err != nil {
return trace.Wrap(err)
}
c.Roles.AllowReview[allow.Where], err = AppendRoleMatchers(c.Roles.AllowReview[allow.Where], allow.Roles, allow.ClaimsToRoles, c.UserState.GetTraits())
if err != nil {
return trace.Wrap(err)
}
return nil
}
// RequestValidator is a helper for validating access requests.
// A user's statically assigned roles are "added" to the
// validator via the push() method, which extracts all the
// relevant rules, performs variable substitutions, and builds
// a set of simple Allow/Deny datastructures. These, in turn,
// are used to validate and expand the access request.
type RequestValidator struct {
logger *slog.Logger
clock clockwork.Clock
opts ValidateRequestOptions
getter RequestValidatorGetter
userState UserState
// autoRequestOnLogin indicates that a Access Request should be created for a user upon
// login. That happens when any of the users's roles has options.request_access "always"
// or "reason".
autoRequestOnLogin bool
// requireReasonForAllRoles indicates that non-empty reason is required for all access
// requests. This happens if any of the user roles has options.request_access "reason".
requireReasonForAllRoles bool
// requiringReasonRoles is a set of role names, which require non-empty reason to be
// specified when requested. The same applies to all requested resources allowed by those
// roles. Such roles are all requestable roles and search_as_roles allowed by a role
// assigned to a user and having spec.allow.request.reason.mode="required" set.
//
// Please note this means, roles having spec.allow.request.reason.mode="required" don't
// necessarily require reason when they are requested themselves. Instead they mark roles
// in spec.allow.request.roles and spec.allow.request.search_as_roles as roles requiring
// reason.
requiringReasonRoles map[string]struct{}
// customPromptRoles is a set of role names, which specifies a custom prompt when requested.
// Such roles are all requestable roles and search_as_roles allowed by a user's role
// which has spec.allow.request.reason.prompt set.
customPromptRoles map[string]string
// reasonPrompts are the prompts to be displayed in the UI for the reason input box. In the
// case of auto-request only the first prompt is displayed for backward compatibility.
reasonPrompts []string
// Used to enforce that the configuration found in the static
// role that defined the search_as_role, is respected.
// An empty map or list means nothing was configured.
kubernetesResource struct {
// allow is a map from the user's allowed search_as_roles to the list of
// kubernetes resource kinds the user is allowed to request with that role.
allow map[string][]types.RequestKubernetesResource
// deny is the list of kubernetes resource kinds the user is explicitly
// denied from requesting.
deny []types.RequestKubernetesResource
}
roles struct {
allowRequest, denyRequest []parse.Matcher
allowSearch, denySearch []string
}
annotations struct {
// allow annotations are not greedy, the role that defines the annotation must allow requesting one
// of the roles that are being requested in order for the annotation to be applied.
allow map[singleAnnotation]annotationMatcher
// deny annotations match greedily, if a user has any role that denies a specific annotation it will
// always be denied.
deny map[singleAnnotation]struct{}
}
thresholdMatchers []struct {
matchers []parse.Matcher
thresholds []types.AccessReviewThreshold
}
suggestedReviewers []string
maxDurationMatchers []struct {
matchers []parse.Matcher
maxDuration time.Duration
}
}
// NewRequestValidator configures a new RequestValidator for the specified user.
func NewRequestValidator(ctx context.Context, clock clockwork.Clock, getter RequestValidatorGetter, username string, opts ...ValidateRequestOption) (RequestValidator, error) {
uls, err := GetUserOrLoginState(ctx, getter, username)
if err != nil {
return RequestValidator{}, trace.Wrap(err)
}
v, err := NewRequestValidatorForUser(ctx, clock, getter, uls, opts...)
if err != nil {
return RequestValidator{}, trace.Wrap(err)
}
return v, nil
}
func NewRequestValidatorForUser(ctx context.Context, clock clockwork.Clock, getter RequestValidatorGetter, user UserState, opts ...ValidateRequestOption) (RequestValidator, error) {
m := RequestValidator{
logger: slog.With(teleport.ComponentKey, "request.validator"),
clock: clock,
getter: getter,
userState: user,
requiringReasonRoles: make(map[string]struct{}),
customPromptRoles: make(map[string]string),
}
for _, opt := range opts {
opt(&m.opts)
}
if m.opts.expandVars {
// validation process for incoming access requests requires
// generating system annotations to be attached to the request
// before it is inserted into the backend.
m.annotations.allow = make(map[singleAnnotation]annotationMatcher)
m.annotations.deny = make(map[singleAnnotation]struct{})
}
m.kubernetesResource.allow = make(map[string][]types.RequestKubernetesResource)
// load all statically assigned roles for the user and
// use them to build our validation state.
for _, roleName := range m.userState.GetRoles() {
role, err := m.getter.GetRole(ctx, roleName)
if err != nil {
return RequestValidator{}, trace.Wrap(err)
}
if err := m.push(ctx, role); err != nil {
return RequestValidator{}, trace.Wrap(err)
}
}
// we retain a fixed order of global reason prompts for determinism of auto-request
// backward compatibility (only the first prompt is displayed for auto-requests)
slices.Sort(m.reasonPrompts)
return m, nil
}
func (m *RequestValidator) roleTemplateContext() RoleTemplateContext {
return RoleTemplateContext{
Username: m.userState.GetName(),
Traits: m.userState.GetTraits(),
}
}
// validate validates an access request and potentially modifies it depending on what the validator
// options were configured in the requestValidator.
//
// When requestValidator.opts.expandVars is true, it expands wildcard requests, setting their role
// list to include all roles the user is allowed to request. Expansion should be performed before
// an access request is initially placed in the backend.
//
// When requestValidator.opts.expandVars is true and req.GetDryRun() is true, it adds expanded
// dry-run enrichment data to the request.
func (m *RequestValidator) validate(ctx context.Context, req types.AccessRequest, identity tlsca.Identity) error {
if m.userState.GetName() != req.GetUser() {
return trace.BadParameter("request validator configured for different user (this is a bug)")
}
if !req.GetState().IsPromoted() && req.GetPromotedAccessListTitle() != "" {
return trace.BadParameter("only promoted requests can set the promoted access list title")
}
// TODO(kiosion): As part of Reviewer changes for long-term requests, roles, expiry, maxDur should not be allowed to be set.
// check for "wildcard request" (`roles=*`). wildcard requests
// need to be expanded into a list consisting of all existing roles
// that the user does not hold and is allowed to request.
if r := req.GetRoles(); len(r) == 1 && r[0] == types.Wildcard {
if !req.GetState().IsPending() {
// expansion is only permitted in pending requests. once resolved,
// a request's role list must be immutable.
return trace.BadParameter("wildcard requests are not permitted in state %s", req.GetState())
}
if !m.opts.expandVars {
// teleport always validates new incoming pending access requests
// with ExpandVars(true). after that, it should be impossible to
// add new values to the role list.
return trace.BadParameter("unexpected wildcard request (this is a bug)")
}
requestable, err := m.getRequestableRoles(ctx, identity, nil, "")
if err != nil {
return trace.Wrap(err)
}
if len(requestable) == 0 {
return trace.BadParameter("no requestable roles, please verify static RBAC configuration")
}
req.SetRoles(requestable)
}
enrichment := &types.AccessRequestDryRunEnrichment{
ReasonMode: types.RequestReasonModeOptional,
}
allRequestedResources := req.GetAllRequestedResourceIDs()
// populate the custom reason prompts: both from global prompts (spec.options.request_prompt) and
// from role/resource specific prompts (spec.allow.request.reason.prompt)
if err := m.populateCustomReasonPrompts(ctx, req.GetRoles(), allRequestedResources); err != nil {
return trace.Wrap(err)
}
// retain a deterministic order of reason prompts
slices.Sort(m.reasonPrompts)
enrichment.ReasonPrompts = m.reasonPrompts
switch {
// for dry-run, store the reason requirement in the enrichment data
case req.GetDryRun():
required, _, err := m.isReasonRequired(ctx, req.GetRoles(), allRequestedResources)
if err != nil {
return trace.Wrap(err)
}
if required {
enrichment.ReasonMode = types.RequestReasonModeRequired
}
// for no dry-run and no reason provided, fail if the reason is required
case len(strings.TrimSpace(req.GetRequestReason())) == 0:
required, explanation, err := m.isReasonRequired(ctx, req.GetRoles(), allRequestedResources)
if err != nil {
return trace.Wrap(err)
}
if required {
promptString := ""
if len(m.reasonPrompts) > 0 {
promptString = "\n" + strings.Join(m.reasonPrompts, "\n")
}
return trace.BadParameter("%s%s", explanation, promptString)
}
}
// verify that all requested roles are permissible
for _, roleName := range req.GetRoles() {
if len(allRequestedResources) > 0 {
if !m.canSearchAsRole(roleName) {
// Roles are normally determined automatically for resource
// access requests, this role must have been explicitly
// requested, or a new deny rule has since been added.
return trace.BadParameter("user %q %s %q", req.GetUser(), CannotRequestRole, roleName)
}
} else {
if !m.CanRequestRole(roleName) {
return trace.BadParameter("user %q %s %q", req.GetUser(), CannotRequestRole, roleName)
}
}
}
// Verify that each requested role allows requesting every requested kube resource kind.
if len(allRequestedResources) > 0 && len(req.GetRoles()) > 0 {
// If there were pruned roles, then the request will be rejected.
// A pruned role meant that role did not allow requesting to all of requested kube resource.
prunedRoles, mappedRequestedRolesToAllowedKinds := m.pruneRequestedRolesNotMatchingKubernetesResourceKinds(allRequestedResources, req.GetRoles())
if len(prunedRoles) != len(req.GetRoles()) {
return getInvalidKubeKindAccessRequestsError(mappedRequestedRolesToAllowedKinds, true /* requestedRoles */)
}
}
if m.opts.expandVars {
// deduplicate requested resource IDs
deduplicateRequestedResources(req)
// In addition to capping the maximum number of resources in a single request,
// we also need to ensure that the sum of the resource IDs in the request doesn't
// get too big.
if err := validateRequestedResourcesLength(req); err != nil {
return trace.Wrap(err)
}
// determine the roles which should be requested for a resource access
// request, and write them to the request
if err := m.setRolesForResourceRequest(ctx, req); err != nil {
return trace.Wrap(err)
}
// build the threshold array and role-threshold-mapping. the rtm encodes the
// relationship between a role, and the thresholds which must pass in order
// for that role to be considered approved. when building the validator we
// recorded the relationship between the various allow matchers and their associated
// threshold groups.
rtm := make(map[string]types.ThresholdIndexSets)
var tc thresholdCollector
for _, role := range req.GetRoles() {
sets, err := m.collectSetsForRole(&tc, role)
if err != nil {
return trace.Wrap(err)
}
rtm[role] = types.ThresholdIndexSets{
Sets: sets,
}
}
req.SetThresholds(tc.Thresholds)
req.SetRoleThresholdMapping(rtm)
// incoming requests must have system annotations attached
// before being inserted into the backend. this is how the
// RBAC system propagates sideband information to plugins.
systemAnnotations, err := m.systemAnnotations(req)
if err != nil {
return trace.Wrap(err)
}
req.SetSystemAnnotations(systemAnnotations)
// if no suggested reviewers were provided by the user, then
// use the defaults suggested by the user's static roles.
if len(req.GetSuggestedReviewers()) == 0 {
req.SetSuggestedReviewers(apiutils.Deduplicate(m.suggestedReviewers))
}
// Pin the time to the current time to prevent time drift.
now := m.clock.Now().UTC()
// TODO(kiosion): The following logic shouldn't be relevant for long-term requests, post-Reviewer-changes.
// Calculate the expiration time of the elevated certificate that will
// be issued if the Access Request is approved.
sessionTTL, err := m.sessionTTL(ctx, identity, req, now)
if err != nil {
return trace.Wrap(err)
}
maxDuration, err := m.calculateMaxAccessDuration(req, sessionTTL)
if err != nil {
return trace.Wrap(err)
}
// If the maxDuration flag is set, consider it instead of only using the session TTL.
var maxAccessDuration time.Duration
if maxDuration > 0 {
req.SetSessionTLL(now.Add(min(sessionTTL, maxDuration)))
maxAccessDuration = maxDuration
} else {
req.SetSessionTLL(now.Add(sessionTTL))
maxAccessDuration = sessionTTL
}
// This is the final adjusted access expiry where both max duration
// and session TTL were taken into consideration.
accessExpiry := now.Add(maxAccessDuration)
// Adjusted max access duration is equal to the access expiry time.
req.SetMaxDuration(accessExpiry)
// Setting access expiry before calling `calculatePendingRequestTTL`
// matters since the func relies on this adjusted expiry.
req.SetAccessExpiry(accessExpiry)
// Calculate the expiration time of the Access Request (how long it
// will await approval).
requestTTL, err := m.calculatePendingRequestTTL(req, now)
if err != nil {
return trace.Wrap(err)
}
req.SetExpiry(now.Add(requestTTL))
if req.GetAssumeStartTime() != nil {
assumeStartTime := *req.GetAssumeStartTime()
if err := types.ValidateAssumeStartTime(assumeStartTime, accessExpiry, req.GetCreationTime()); err != nil {
return trace.Wrap(err)
}
}
}
if req.GetDryRun() {
req.SetDryRunEnrichment(enrichment)
}
return nil
}
// validateRequestedResourcesLength checks the sum length of requested resources when serialized
// against the maximum allowed length.
func validateRequestedResourcesLength(req types.AccessRequest) error {
resourcesLen := 0
requestedResourceAccessIDs := req.GetRequestedResourceAccessIDs()
requestedResourceIDs := req.GetRequestedResourceIDs()
if len(requestedResourceAccessIDs) > 0 {
str, err := types.ResourceAccessIDsToString(requestedResourceAccessIDs)
if err != nil {
return trace.Wrap(err)
}
resourcesLen += len(str)
}
if len(requestedResourceIDs) > 0 {
str, err := types.ResourceIDsToString(requestedResourceIDs)
if err != nil {
return trace.Wrap(err)
}
resourcesLen += len(str)
}
if resourcesLen > maxResourcesLength {
return trace.BadParameter("access request exceeds maximum length: try reducing the number of resources in the request")
}
return nil
}
// deduplicateRequestedResources deduplicates both requestedResourceIDs and requestedResourceAccessIDs,
// ensuring there are no duplicate entries in either or across both. Constrained resources take precedence.
func deduplicateRequestedResources(req types.AccessRequest) {
var deduplicatedResourceIDs []types.ResourceID
var deduplicatedResourceAccessIDs []types.ResourceAccessID
seen := make(map[string]struct{})
// first deduplicate requestedResourceAccessIDs so resources carrying constraints take precedence
for _, pair := range req.GetRequestedResourceAccessIDs() {
idStr := types.ResourceIDToString(pair.Id)
if _, isDuplicate := seen[idStr]; isDuplicate {
continue
}
seen[idStr] = struct{}{}
deduplicatedResourceAccessIDs = append(deduplicatedResourceAccessIDs, pair)
}
for _, resource := range req.GetRequestedResourceIDs() {
idStr := types.ResourceIDToString(resource)
if _, isDuplicate := seen[idStr]; isDuplicate {
continue
}
seen[idStr] = struct{}{}
deduplicatedResourceIDs = append(deduplicatedResourceIDs, resource)
}
req.SetRequestedResourceIDs(deduplicatedResourceIDs)
req.SetRequestedResourceAccessIDs(deduplicatedResourceAccessIDs)
}
// isReasonRequired checks if the reason is required for the given roles and resource IDs.
func (v *RequestValidator) isReasonRequired(ctx context.Context, requestedRoles []string, requestedResourceAccessIDs []types.ResourceAccessID) (required bool, explanation string, err error) {
if v.requireReasonForAllRoles {
return true, "request reason must be specified (required request_access option in one of the roles)", nil
}
allApplicableRoles, err := v.getAllApplicableRoles(ctx, requestedRoles, requestedResourceAccessIDs)
if err != nil {
return false, "", trace.Wrap(err)
}
for _, r := range allApplicableRoles {
if _, ok := v.requiringReasonRoles[r]; ok {
return true, fmt.Sprintf("request reason must be specified (required for role %q)", r), nil
}
}
return false, "", nil
}
func (v *RequestValidator) populateCustomReasonPrompts(ctx context.Context, requestedRoles []string, requestedResourceAccessIDs []types.ResourceAccessID) error {
allApplicableRoles, err := v.getAllApplicableRoles(ctx, requestedRoles, requestedResourceAccessIDs)
if err != nil {
return trace.Wrap(err)
}
for _, r := range allApplicableRoles {
customPrompt, ok := v.customPromptRoles[r]
if ok && !slices.Contains(v.reasonPrompts, customPrompt) {
v.reasonPrompts = append(v.reasonPrompts, customPrompt)
}
}
return nil
}
// getAllApplicableRoles returns the combined roles for the given roles and resource IDs (search_as_roles)
func (v *RequestValidator) getAllApplicableRoles(ctx context.Context, requestedRoles []string, requestedResourceAccessIDs []types.ResourceAccessID) (allApplicableRoles []string, err error) {
allApplicableRoles = requestedRoles
if len(requestedResourceAccessIDs) > 0 {
// Do not provide loginHint. We want all matching search_as_roles for those resources.
roles, err := v.applicableSearchAsRoles(ctx, requestedResourceAccessIDs, "")
if err != nil {
return nil, trace.Wrap(err)
}
if len(allApplicableRoles) == 0 {
allApplicableRoles = roles
} else {
allApplicableRoles = append(allApplicableRoles, roles...)
}
}
return allApplicableRoles, nil
}
// calculateMaxAccessDuration calculates the maximum time for the access request.
// The max duration time is the minimum of the max_duration time set on the request
// and the max_duration time set on the request role.
func (m *RequestValidator) calculateMaxAccessDuration(req types.AccessRequest, sessionTTL time.Duration) (time.Duration, error) {
// Check if the maxDuration time is set.
maxDurationTime := req.GetMaxDuration()
maxDuration := maxDurationTime.Sub(req.GetCreationTime())
// For dry run requests, use the maximum possible duration.
// This prevents the time drift that can occur as the value is set on the client side.
if req.GetDryRun() {
maxDuration = MaxAccessDuration
// maxDuration may end up < 0 even if maxDurationTime is set
} else if !maxDurationTime.IsZero() && maxDuration < 0 {
return 0, trace.BadParameter("invalid maxDuration: must be greater than creation time")
}
if maxDuration > MaxAccessDuration {
return 0, trace.BadParameter("max_duration must be less than or equal to %v", MaxAccessDuration)
}
var minAdjDuration time.Duration
// Adjust the expiration time if the max_duration value is set on the request role.
for _, roleName := range req.GetRoles() {
maxDurationForRole := m.maxDurationForRole(roleName)
if minAdjDuration == 0 || maxDurationForRole < minAdjDuration {
minAdjDuration = maxDurationForRole
}
}
if !maxDurationTime.IsZero() && maxDuration < minAdjDuration {
minAdjDuration = maxDuration
}
// minAdjDuration can end up being 0, if no role has a
// field `max_duration` defined.
// In this case, return the smaller value between the sessionTTL
// and the requested max duration.
if minAdjDuration == 0 && maxDuration < sessionTTL {
return maxDuration, nil
}
return minAdjDuration, nil
}
func (m *RequestValidator) maxDurationForRole(roleName string) time.Duration {
var maxDurationForRole time.Duration
for _, tms := range m.maxDurationMatchers {
for _, matcher := range tms.matchers {
if matcher.Match(roleName) {
if tms.maxDuration > maxDurationForRole {
maxDurationForRole = tms.maxDuration
}
}
}
}
return maxDurationForRole
}
// calculatePendingRequestTTL calculates the TTL of the Access Request (how long it will await
// approval). request TTL is capped to the smaller value between the const requestTTL and the
// access request access expiry.
func (m *RequestValidator) calculatePendingRequestTTL(r types.AccessRequest, now time.Time) (time.Duration, error) {
accessExpiryTTL := r.GetAccessExpiry().Sub(now)
// If no expiration provided, use default.
expiry := r.Expiry()
if expiry.IsZero() {
// Guard against the default expiry being greater than access expiry.
if requestTTL < accessExpiryTTL {
expiry = now.Add(requestTTL)
} else {
expiry = now.Add(accessExpiryTTL)
}
}
if expiry.Before(now) {
return 0, trace.BadParameter("invalid request TTL: Access Request can not be created in the past")
}
// Before returning the TTL, validate that the value requested was smaller
// than the maximum value allowed. Used to return a sensible error to the
// user.
requestedTTL := expiry.Sub(now)
if !r.Expiry().IsZero() {
if requestedTTL > requestTTL {
return 0, trace.BadParameter("invalid request TTL: %v greater than maximum allowed (%v)", requestedTTL, requestTTL)
}
if requestedTTL > accessExpiryTTL {
return 0, trace.BadParameter("invalid request TTL: %v greater than maximum allowed (%v)", requestedTTL, accessExpiryTTL)
}
}
return requestedTTL, nil
}
// sessionTTL calculates the TTL of the elevated certificate that will be issued
// if the Access Request is approved.
func (m *RequestValidator) sessionTTL(ctx context.Context, identity tlsca.Identity, r types.AccessRequest, now time.Time) (time.Duration, error) {
ttl, err := m.truncateTTL(ctx, identity, r.GetAccessExpiry(), r.GetRoles(), now)
if err != nil {
return 0, trace.BadParameter("invalid session TTL: %v", err)
}
// Before returning the TTL, validate that the value requested was smaller
// than the maximum value allowed. Used to return a sensible error to the
// user.
requestedTTL := r.GetAccessExpiry().Sub(now)
if !r.GetAccessExpiry().IsZero() && requestedTTL > ttl {
return 0, trace.BadParameter("invalid session TTL: %v greater than maximum allowed (%v)", requestedTTL, ttl)
}
return ttl, nil
}
// truncateTTL will truncate given expiration by identity expiration and
// shortest session TTL of any role.
func (m *RequestValidator) truncateTTL(ctx context.Context, identity tlsca.Identity, expiry time.Time, roles []string, now time.Time) (time.Duration, error) {
ttl := apidefaults.MaxCertDuration
// Reduce by remaining TTL on requesting certificate (identity).
identityTTL := identity.Expires.Sub(now)
if identityTTL > 0 && identityTTL < ttl {
ttl = identityTTL
}
// Reduce TTL further if expiration time requested is shorter than that
// identity.
expiryTTL := expiry.Sub(now)
if expiryTTL > 0 && expiryTTL < ttl {
ttl = expiryTTL
}
// Loop over the roles requested by the user and reduce certificate TTL
// further. Follow the typical Teleport RBAC pattern of strictest setting
// wins.
for _, roleName := range roles {
role, err := m.getter.GetRole(ctx, roleName)
if err != nil {
return 0, trace.Wrap(err)
}
roleTTL := time.Duration(role.GetOptions().MaxSessionTTL)
if roleTTL > 0 && roleTTL < ttl {
ttl = roleTTL
}
}
return ttl, nil
}
// getResourceViewingRoles gets the subset of the user's roles that could be used
// to view resources (i.e., base roles + search as roles).
func (m *RequestValidator) getResourceViewingRoles() []string {
roles := slices.Clone(m.userState.GetRoles())
for _, role := range m.roles.allowSearch {
if m.canSearchAsRole(role) {
roles = append(roles, role)
}
}
return apiutils.Deduplicate(roles)
}
// getRequestableRoles gets the list of all existent roles which the user is
// able to request. This operation is expensive since it loads all existent
// roles to determine the role list. Prefer calling CanRequestRole
// when checking against a known role list. If resource IDs or a login hints
// are provided, roles will be filtered to only include those that would
// allow access to the given resource with the given login.
func (m *RequestValidator) getRequestableRoles(ctx context.Context, identity tlsca.Identity, resourceAccessIDs []types.ResourceAccessID, loginHint string) ([]string, error) {
allRoles, err := m.getter.GetRoles(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// For fetching the underling resources, we can safely discard additional info carried on the ResourceAccessID.
underlyingResources, err := m.getUnderlyingResourcesByResourceIDs(ctx, types.RiskyExtractResourceIDs(resourceAccessIDs))
if err != nil {
return nil, trace.Wrap(err)
}
cluster, err := m.getter.GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := NewAccessChecker(&AccessInfo{
Roles: m.getResourceViewingRoles(),
Traits: m.userState.GetTraits(),
Username: m.userState.GetName(),
AllowedResourceAccessIDs: identity.AllowedResourceAccessIDs,
DelegationSessionID: identity.DelegationSessionID,
}, cluster.GetClusterName(), m.getter)
if err != nil {
return nil, trace.Wrap(err)
}
// Filter out resources the user requested but doesn't have access to and
// pair each remaining resource with matchers for any requested constraints.
filteredResources := make([]types.ResourceWithLabels, 0, len(underlyingResources))
constraintMatchers := make([][]RoleMatcher, 0, len(underlyingResources))
for _, resource := range underlyingResources {
if err := accessChecker.CheckAccess(resource, AccessState{MFAVerified: true}); err == nil {
matchers, err := BuildResourceConstraintMatchers(resourceAccessIDs, resource)
if err != nil {
return nil, trace.Wrap(err)
}
filteredResources = append(filteredResources, resource)
constraintMatchers = append(constraintMatchers, matchers)
}
}
var expanded []string
for _, role := range allRoles {
n := role.GetName()
if slices.Contains(m.userState.GetRoles(), n) || !m.CanRequestRole(n) {
continue
}
roleAllowsAccess := true
for i, resource := range filteredResources {
access, err := m.roleAllowsResource(role, resource, loginHint, constraintMatchers[i]...)
if err != nil {
return nil, trace.Wrap(err)
}
if !access {
roleAllowsAccess = false
}
}
// user does not currently hold this role, and is allowed to request it.
if roleAllowsAccess {
expanded = append(expanded, n)
}
}
return expanded, nil
}
// setAllowRequestKubeResourceLookup goes through each search as roles and sets it with the allowed roles.
// Multiple allow request.kubernetes_resources found for a role will be merged, except when an empty configuration
// is encountered. In this case, empty configuration will override configured request field
// (which results in allowing anything).
func setAllowRequestKubeResourceLookup(allowKubernetesResources []types.RequestKubernetesResource, searchAsRoles []string, lookup map[string][]types.RequestKubernetesResource) {
if len(allowKubernetesResources) == 0 {
// Empty configuration overrides any configured request.kubernetes_resources field.
for _, searchAsRoles := range searchAsRoles {
lookup[searchAsRoles] = []types.RequestKubernetesResource{}
}
return
}
for _, searchAsRole := range searchAsRoles {
currentAllowedResources, exists := lookup[searchAsRole]
if exists && len(currentAllowedResources) == 0 {
// Already allowed to access all kube resource kinds.
continue
}
lookup[searchAsRole] = append(currentAllowedResources, allowKubernetesResources...)
}
}
// push compiles a role's configuration into the request validator.
// All of the requesting user's statically assigned roles must be pushed
// before validation begins.
func (m *RequestValidator) push(ctx context.Context, role types.Role) error {
var err error
m.requireReasonForAllRoles = m.requireReasonForAllRoles || role.GetOptions().RequestAccess.RequireReason()
m.autoRequestOnLogin = m.autoRequestOnLogin || role.GetOptions().RequestAccess.ShouldAutoRequest()
reasonPrompt := strings.TrimSpace(role.GetOptions().RequestPrompt)
if len(reasonPrompt) > 0 && !slices.Contains(m.reasonPrompts, reasonPrompt) {
m.reasonPrompts = append(m.reasonPrompts, reasonPrompt)
}
allow, deny := role.GetAccessRequestConditions(types.Allow), role.GetAccessRequestConditions(types.Deny)
if allow.Reason != nil {
if allow.Reason.Mode.Required() {
for _, r := range allow.Roles {
m.requiringReasonRoles[r] = struct{}{}
}
for _, r := range allow.SearchAsRoles {
m.requiringReasonRoles[r] = struct{}{}
}
}
customPrompt := strings.TrimSpace(allow.Reason.Prompt)
if len(customPrompt) > 0 {
for _, r := range allow.Roles {
m.customPromptRoles[r] = customPrompt
}
for _, r := range allow.SearchAsRoles {
m.customPromptRoles[r] = customPrompt
}
}
}
// NOTE: Not using allow.KubernetesResources as we need to map older roles to new values.
setAllowRequestKubeResourceLookup(role.GetRequestKubernetesResources(types.Allow), allow.SearchAsRoles, m.kubernetesResource.allow)
if deniedKubeResources := role.GetRequestKubernetesResources(types.Deny); len(deniedKubeResources) > 0 {
m.kubernetesResource.deny = append(m.kubernetesResource.deny, deniedKubeResources...)
}
m.roles.denyRequest, err = AppendRoleMatchers(m.roles.denyRequest, deny.Roles, deny.ClaimsToRoles, m.userState.GetTraits())
if err != nil {
return trace.Wrap(err)
}
// record what will be the starting index of the allow and deny matchers for this role, if it applies any.
astart := len(m.roles.allowRequest)
m.roles.allowRequest, err = AppendRoleMatchers(m.roles.allowRequest, allow.Roles, allow.ClaimsToRoles, m.userState.GetTraits())
if err != nil {
return trace.Wrap(err)
}
m.roles.allowSearch = apiutils.Deduplicate(append(m.roles.allowSearch, allow.SearchAsRoles...))
m.roles.denySearch = apiutils.Deduplicate(append(m.roles.denySearch, deny.SearchAsRoles...))
if m.opts.expandVars {
// if this role added additional allow matchers, then we need to record the relationship
// between its matchers and its thresholds. This information is used later to calculate
// the rtm and threshold list.
newAllowRequestMatchers := m.roles.allowRequest[astart:]
newAllowSearchMatchers := literalMatchers(allow.SearchAsRoles)
allNewAllowMatchers := make([]parse.Matcher, 0, len(newAllowRequestMatchers)+len(newAllowSearchMatchers))
allNewAllowMatchers = append(allNewAllowMatchers, newAllowRequestMatchers...)
allNewAllowMatchers = append(allNewAllowMatchers, newAllowSearchMatchers...)
if len(allNewAllowMatchers) > 0 {
m.thresholdMatchers = append(m.thresholdMatchers, struct {
matchers []parse.Matcher
thresholds []types.AccessReviewThreshold
}{
matchers: allNewAllowMatchers,
thresholds: allow.Thresholds,
})
}
if allow.MaxDuration != 0 {
m.maxDurationMatchers = append(m.maxDurationMatchers, struct {
matchers []parse.Matcher
maxDuration time.Duration
}{
matchers: allNewAllowMatchers,
maxDuration: allow.MaxDuration.Duration(),
})
}
// validation process for incoming access requests requires
// generating system annotations to be attached to the request
// before it is inserted into the backend.
m.insertAllowedAnnotations(ctx, allow, newAllowRequestMatchers, newAllowSearchMatchers)
m.insertDeniedAnnotations(ctx, deny)
m.suggestedReviewers = append(m.suggestedReviewers, allow.SuggestedReviewers...)
}
return nil
}
// setRolesForResourceRequest determines if the given access request is
// resource-based, and if so, it determines which underlying roles are necessary
// and adds them to the request.
func (m *RequestValidator) setRolesForResourceRequest(ctx context.Context, req types.AccessRequest) error {
if !m.opts.expandVars {
// Don't set the roles if expandVars is not set, they have probably
// already been set and we are just validating the request.
return nil
}
if req.GetRequestKind().IsLongTerm() {
// Don't set roles on LongTerm requests; they are only allowed
// to be search-based resource requests.
return nil
}
if len(req.GetAllRequestedResourceIDs()) == 0 {
// This is not a resource request.
return nil
}
if len(req.GetRoles()) > 0 {
// Roles were explicitly requested, don't change them.
return nil
}
rolesToRequest, err := m.applicableSearchAsRoles(ctx, req.GetAllRequestedResourceIDs(), req.GetLoginHint())
if err != nil {
return trace.Wrap(err)
}
req.SetRoles(rolesToRequest)
return nil
}
// requestResourcesToStrings formats the resource list as <kind>.<apiGroup>.
// Removes wildcards if any.
func requestResourcesToStrings(resources, denied []types.RequestKubernetesResource) []string {
strs := make([]string, 0, len(resources))
for _, resource := range resources {
str := resource.Kind
if resource.APIGroup != "" {
str += "." + resource.APIGroup
}
if resource.Kind == types.Wildcard && len(denied) > 0 {
str += "(- " + strings.Join(requestResourcesToStrings(denied, nil), ", ") + ")"
}
strs = append(strs, str)
}
return strs
}
// pruneRequestedRolesNotMatchingKubernetesResourceKinds will filter out the kubernetes kinds from the requested resource IDs (kube_cluster and its subresources)
// disregarding whether it's leaf or root cluster request, and for each requested role, ensures that all requested kube resource kind are allowed by the role.
// Roles not matching with every kind requested, will be pruned from the requested roles.
//
// Returns pruned roles, and a map of requested roles with allowed kinds (with denied applied), used to help aid user in case a request gets rejected,
// lets user know which kinds are allowed for each requested roles.
func (m *RequestValidator) pruneRequestedRolesNotMatchingKubernetesResourceKinds(requestedResourceAccessIDs []types.ResourceAccessID, requestedRoles []string) ([]string, map[string][]string) {
// Filter for the kube_cluster and its subresource kinds.
requestedKubeKinds := map[gk]struct{}{}
for _, resource := range requestedResourceAccessIDs {
resourceID := resource.GetResourceID()
if resourceID.Kind == types.KindKubernetesCluster {
requestedKubeKinds[gk{kind: types.KindKubernetesCluster}] = struct{}{}
} else if slices.Contains(types.KubernetesResourcesKinds, resourceID.Kind) || strings.HasPrefix(resourceID.Kind, types.AccessRequestPrefixKindKube) {
requestedKubeKinds[normalizeKubernetesKind(resourceID.Kind)] = struct{}{}
}
}
if len(requestedKubeKinds) == 0 {
return requestedRoles, nil
}
goodRoles := map[string]struct{}{}
mappedRequestedRolesToAllowedKinds := map[string][]string{}
for _, requestedRoleName := range requestedRoles {
allowedKinds, deniedKinds := m.kubernetesResource.allow[requestedRoleName], m.kubernetesResource.deny
// If there is nothing in allowed nor deny, everything is allowed.
if len(allowedKinds) == 0 && len(deniedKinds) == 0 {
goodRoles[requestedRoleName] = struct{}{}
continue
}
// All supported kube kinds are allowed when there was nothing configured.
if len(allowedKinds) == 0 {
allowedKinds = append(allowedKinds,
types.RequestKubernetesResource{Kind: types.Wildcard, APIGroup: types.Wildcard},
)
// If there is nothing in deny, also include kube_cluster.
if len(deniedKinds) == 0 {
allowedKinds = append(allowedKinds,
types.RequestKubernetesResource{Kind: types.KindKubernetesCluster},
)
}
}
allowedKinds = slices.DeleteFunc(allowedKinds, func(in types.RequestKubernetesResource) bool {
for _, elem := range deniedKinds {
if matchRequestKubernetesResources(gk{group: in.APIGroup, kind: in.Kind}, elem, types.Allow) {
return true
}
}
return false
})
// TODO(@creack): Consider removing this. We shouldn't disclose to the user what they could request when getting an access denied error.
// Keeping existing behavior for now.
mappedRequestedRolesToAllowedKinds[requestedRoleName] = requestResourcesToStrings(allowedKinds, deniedKinds)
// If we have any requested kinds that is either not allowed or that is denied, reject the role.
// TODO(@creack): Reconsider this, we may want to allow some kinds and deny others.
// Keeping existing behavior for now.
filteredAllowedKinds := make([]types.RequestKubernetesResource, 0, len(requestedKubeKinds))
for requestedKubeKind := range requestedKubeKinds {
for _, k := range allowedKinds {
if matchRequestKubernetesResources(requestedKubeKind, k, types.Allow) {
filteredAllowedKinds = append(filteredAllowedKinds, types.RequestKubernetesResource{Kind: requestedKubeKind.kind, APIGroup: requestedKubeKind.group})
break
}
}
}
if len(filteredAllowedKinds) != len(requestedKubeKinds) {
// If we don't have as many allowed kinds as request, we reject the role.
continue
}
// If there is something to deny, make sure we reject 'namespace' and 'kube_cluster', as it would grant access to everything.
for requestedKubeKind := range requestedKubeKinds {
for _, k := range deniedKinds {
if requestedKubeKind.kind == types.KindKubernetesCluster || requestedKubeKind.kind == "namespaces" {
// We have a deny entry and the request is for a kube_cluster or namespaces, reject.
return nil, mappedRequestedRolesToAllowedKinds
}
if matchRequestKubernetesResources(requestedKubeKind, k, types.Deny) {
// If we have any requested kinds that is denied, reject all roles.
return nil, mappedRequestedRolesToAllowedKinds
}
}
}
goodRoles[requestedRoleName] = struct{}{}
}
return slices.Collect(maps.Keys(goodRoles)), mappedRequestedRolesToAllowedKinds
}
// thresholdCollector is a helper that assembles the Thresholds array for a request.
// the push() method is used to insert groups of related thresholds and calculate their
// corresponding index set.
type thresholdCollector struct {
Thresholds []types.AccessReviewThreshold
}
// push pushes a set of related thresholds and returns the associated indexes. each set of indexes represents
// an "or" operator, indicating that one of the referenced thresholds must reach its approval condition in order
// for the set as a whole to be considered approved.
func (c *thresholdCollector) push(s []types.AccessReviewThreshold) ([]uint32, error) {
if len(s) == 0 {
// empty threshold sets are equivalent to the default threshold
s = []types.AccessReviewThreshold{
{
Name: "default",
Approve: 1,
Deny: 1,
},
}
}
var indexes []uint32
for _, t := range s {
tid, err := c.pushThreshold(t)
if err != nil {
return nil, trace.Wrap(err)
}
indexes = append(indexes, tid)
}
return indexes, nil
}
// pushThreshold pushes a threshold to the main threshold list and returns its index
// as a uint32 for compatibility with grpc types.
func (c *thresholdCollector) pushThreshold(t types.AccessReviewThreshold) (uint32, error) {
// maxThresholds is an arbitrary large number that serves as a guard against
// odd errors due to casting between int and uint32. This is probably unnecessary
// since we'd likely hit other limitations *well* before wrapping became a concern,
// but its best to have explicit guard rails.
const maxThresholds = 4096
// don't bother double-storing equivalent thresholds
for i, threshold := range c.Thresholds {
if t.IsEqual(&threshold) {
return uint32(i), nil
}
}
if len(c.Thresholds) >= maxThresholds {
return 0, trace.LimitExceeded("max review thresholds exceeded (max=%d)", maxThresholds)
}
c.Thresholds = append(c.Thresholds, t)
return uint32(len(c.Thresholds) - 1), nil
}
// CanRequestRole checks if a given role can be requested.
func (m *RequestValidator) CanRequestRole(name string) bool {
for _, deny := range m.roles.denyRequest {
if deny.Match(name) {
return false
}
}
for _, allow := range m.roles.allowRequest {
if allow.Match(name) {
return true
}
}
return false
}
// canSearchAsRole check if a given role can be requested through a search-based
// access request
func (m *RequestValidator) canSearchAsRole(name string) bool {
if slices.Contains(m.roles.denySearch, name) {
return false
}
for _, deny := range m.roles.denyRequest {
if deny.Match(name) {
return false
}
}
return slices.Contains(m.roles.allowSearch, name)
}
// collectSetsForRole collects the threshold index sets which describe the various groups of
// thresholds which must pass in order for a request for the given role to be approved.
func (m *RequestValidator) collectSetsForRole(c *thresholdCollector, role string) ([]types.ThresholdIndexSet, error) {
var sets []types.ThresholdIndexSet
Outer:
for _, tms := range m.thresholdMatchers {
for _, matcher := range tms.matchers {
if matcher.Match(role) {
set, err := c.push(tms.thresholds)
if err != nil {
return nil, trace.Wrap(err)
}
sets = append(sets, types.ThresholdIndexSet{
Indexes: set,
})
continue Outer
}
}
}
if len(sets) == 0 {
// this should never happen since every allow directive is associated with at least one
// threshold, and this operation happens after requested roles have been validated to match at
// least one allow directive.
return nil, trace.BadParameter("role %q matches no threshold sets (this is a bug)", role)
}
return sets, nil
}
// singleAnnotation holds a single annotation key/value pair. The value must already have been expanded with
// ApplyValueTraits.
type singleAnnotation struct {
key, value string
}
// annotationsMatcher holds a set of role matchers used to decide if an annotations should be added to an
// access request when one of the requested roles matches.
type annotationMatcher struct {
roleRequestMatchers []parse.Matcher
resourceRequestMatchers []parse.Matcher
}
// matchesRequest returns true if either:
// - req is a role access request and one of [m.roleRequestMatchers] matches one of the requested roles
// - req is a resource access request and one of [m.resourceRequestMatchers] matches one of the requested roles
func (m *annotationMatcher) matchesRequest(req types.AccessRequest) bool {
matchers := m.roleRequestMatchers
if len(req.GetAllRequestedResourceIDs()) > 0 {
matchers = m.resourceRequestMatchers
}
for _, matcher := range matchers {
if slices.ContainsFunc(req.GetRoles(), matcher.Match) {
return true
}
}
return false
}
// insertAllowedAnnotations constructs all allowed annotations for a given AccessRequestConditions instance
// from one of the users current roles and adds them to the annotation matchers mapping.
//
// Annotations are only applied to access requests requests when one of the requested roles matches one of the
// role matchers.
func (m *RequestValidator) insertAllowedAnnotations(ctx context.Context, conditions types.AccessRequestConditions, roleRequestMatchers, resourceRequestMatchers []parse.Matcher) {
for annotationKey, annotationValueTemplates := range conditions.Annotations {
// iterate through all new values and expand any
// variable interpolation syntax they contain.
for _, template := range annotationValueTemplates {
expandedValues, err := ApplyValueTraitsWithContext(template, m.roleTemplateContext())
if err != nil {
// skip values that failed variable expansion
m.logger.WarnContext(ctx, "Failed to expand trait template in access request annotation",
"key", annotationKey, "template", template, "error", err)
continue
}
for _, expanded := range expandedValues {
annotation := singleAnnotation{annotationKey, expanded}
matchers := m.annotations.allow[annotation]
matchers.roleRequestMatchers = append(matchers.roleRequestMatchers, roleRequestMatchers...)
matchers.resourceRequestMatchers = append(matchers.resourceRequestMatchers, resourceRequestMatchers...)
m.annotations.allow[annotation] = matchers
}
}
}
}
// insertDeniedAnnotations constructs all denied annotations for a given AccessRequestConditions instance
// from one of the users current roles and adds them to the denied annotations set.
func (m *RequestValidator) insertDeniedAnnotations(ctx context.Context, conditions types.AccessRequestConditions) {
for annotationKey, annotationValueTemplates := range conditions.Annotations {
// iterate through all new values and expand any
// variable interpolation syntax they contain.
for _, template := range annotationValueTemplates {
expandedValues, err := ApplyValueTraitsWithContext(template, m.roleTemplateContext())
if err != nil {
// skip values that failed variable expansion
m.logger.WarnContext(ctx, "Failed to expand trait template in access request annotation",
"key", annotationKey, "template", template, "error", err)
continue
}
for _, expanded := range expandedValues {
annotation := singleAnnotation{annotationKey, expanded}
m.annotations.deny[annotation] = struct{}{}
}
}
}
}
// systemAnnotations calculates the system annotations for a pending
// access request.
func (m *RequestValidator) systemAnnotations(req types.AccessRequest) (map[string][]string, error) {
annotations := make(map[string][]string)
for annotation, allowMatchers := range m.annotations.allow {
if _, denied := m.annotations.deny[annotation]; denied {
// Deny matches are greedy, if any of the users roles denies this annotation it is filtered out.
continue
}
if !allowMatchers.matchesRequest(req) {
// Annotations are filtered out unless this request matches one of the role matchers for this
// annotation.
continue
}
annotations[annotation.key] = append(annotations[annotation.key], annotation.value)
}
// Sort and deduplicate.
for k := range annotations {
slices.Sort(annotations[k])
annotations[k] = slices.Compact(annotations[k])
}
return annotations, nil
}
type ValidateRequestOptions struct {
expandVars bool
}
type ValidateRequestOption func(*ValidateRequestOptions)
// WithExpandVars toggles variable expansion during request validation. Variable expansion includes
// expanding wildcard requests, setting system annotations, finding applicable roles for
// resource-based requests and gathering threshold information. Variable expansion should be run
// by the auth server prior to storing an access request for the first time.
func WithExpandVars(expandVars bool) ValidateRequestOption {
return func(v *ValidateRequestOptions) {
v.expandVars = expandVars
}
}
// ValidateAccessRequestForUser validates an access request against the associated users's
// *statically assigned* roles.
//
// It can modify the request.
//
// If [WithExpandVars] is set to true, it will also expand wildcard requests, setting their role
// list to include all roles the user is allowed to request. Expansion should be performed before
// an access request is initially placed in the backend.
//
// If both [WithExpandVars] is set to true and req.GetDryRun() is true it adds expanded dry-run
// enrichment data in the provided request.
func ValidateAccessRequestForUser(ctx context.Context, clock clockwork.Clock, getter RequestValidatorGetter, req types.AccessRequest, identity tlsca.Identity, opts ...ValidateRequestOption) error {
v, err := NewRequestValidator(ctx, clock, getter, req.GetUser(), opts...)
if err != nil {
return trace.Wrap(err)
}
return trace.Wrap(v.validate(ctx, req, identity))
}
// UnmarshalAccessRequest unmarshals the AccessRequest resource from JSON.
func UnmarshalAccessRequest(data []byte, opts ...MarshalOption) (*types.AccessRequestV3, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var req types.AccessRequestV3
if err := utils.FastUnmarshal(data, &req); err != nil {
return nil, trace.Wrap(err)
}
// Requests written by newer Auths must stay readable.
if err := validateAccessRequest(&req, true); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
req.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
req.SetExpiry(cfg.Expires)
}
return &req, nil
}
// MarshalAccessRequest marshals the AccessRequest resource to JSON.
func MarshalAccessRequest(accessRequest types.AccessRequest, opts ...MarshalOption) ([]byte, error) {
// Writes stay strict; re-persisting a request whose constraints
// this build couldn't decode would overwrite the newer content in
// the backend.
if err := ValidateAccessRequest(accessRequest); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch accessRequest := accessRequest.(type) {
case *types.AccessRequestV3:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, accessRequest))
default:
return nil, trace.BadParameter("unrecognized access request type: %T", accessRequest)
}
}
// MarshalAccessRequestAllowedPromotion marshals the list of access list IDs to JSON.
func MarshalAccessRequestAllowedPromotion(accessListIDs *types.AccessRequestAllowedPromotions) ([]byte, error) {
payload, err := utils.FastMarshal(accessListIDs)
return payload, trace.Wrap(err)
}
// UnmarshalAccessRequestAllowedPromotion unmarshals the list of access list IDs from JSON.
func UnmarshalAccessRequestAllowedPromotion(data []byte) (*types.AccessRequestAllowedPromotions, error) {
var accessListIDs types.AccessRequestAllowedPromotions
if err := utils.FastUnmarshal(data, &accessListIDs); err != nil {
return nil, trace.Wrap(err)
}
return &accessListIDs, nil
}
func getInvalidKubeKindAccessRequestsError(mappedRequestedRolesToAllowedKinds map[string][]string, requestedRoles bool) error {
allowedStr := ""
for roleName, allowedKinds := range mappedRequestedRolesToAllowedKinds {
if len(allowedStr) > 0 {
allowedStr = fmt.Sprintf("%s, %s: %v", allowedStr, roleName, allowedKinds)
} else {
allowedStr = fmt.Sprintf("%s: %v", roleName, allowedKinds)
}
}
requestWord := "requestable"
if requestedRoles {
requestWord = "requested"
}
// This error must be in sync with web UI's RequestCheckout.tsx ("checkSupportForKubeResources").
// Web UI relies on the exact format of this error message to determine what kube kinds are
// supported since web UI does not support all kube resources at this time.
return trace.BadParameter(`%s did not allow requesting to some or all of the requested `+
`Kubernetes resources. allowed kinds for each %s roles: %v`,
InvalidKubernetesKindAccessRequest, requestWord, allowedStr)
}
// pruneResourceRequestRoles takes a list of requested resource IDs and
// a list of candidate roles to request, and returns a "pruned" list of roles.
//
// Candidate roles are *always* pruned when the user is not allowed to
// request the role with all requested resources.
//
// A best-effort attempt is made to prune roles that would not allow
// access to any of the requested resources, this is skipped when any
// resource is in a leaf cluster.
//
// If loginHint is provided, it will attempt to prune the list to a single role.
func (m *RequestValidator) pruneResourceRequestRoles(
ctx context.Context,
requestedResourceAccessIDs []types.ResourceAccessID,
loginHint string,
roles []string,
) ([]string, error) {
if len(requestedResourceAccessIDs) == 0 {
// This is not a resource request, nothing to do
return roles, nil
}
roles, mappedRequestedRolesToAllowedKinds := m.pruneRequestedRolesNotMatchingKubernetesResourceKinds(requestedResourceAccessIDs, roles)
if len(roles) == 0 { // all roles got pruned from not matching every kube requested kind.
return nil, getInvalidKubeKindAccessRequestsError(mappedRequestedRolesToAllowedKinds, false /* requestedRoles */)
}
clusterNameResource, err := m.getter.GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
localClusterName := clusterNameResource.GetClusterName()
for _, resource := range requestedResourceAccessIDs {
resourceID := resource.GetResourceID()
if resourceID.ClusterName != localClusterName {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, `Requested resource is in a foreign cluster, unable to prune roles - All available "search_as_roles" will be requested`,
slog.Any("requested_resources", types.ResourceIDToString(resourceID)),
)
return roles, nil
}
}
allRoles, err := FetchRolesWithContext(roles, m.getter, m.roleTemplateContext())
if err != nil {
return nil, trace.Wrap(err)
}
// For fetching the underling resources, we can safely discard additional info carried on the ResourceAccessID.
requestedResourceIDs := types.RiskyExtractResourceIDs(requestedResourceAccessIDs)
underlyingResources, err := m.getUnderlyingResourcesByResourceIDs(ctx, requestedResourceIDs)
if err != nil {
return nil, trace.Wrap(err)
}
necessaryRoles := make(map[string]struct{})
for _, resource := range underlyingResources {
var constraints *types.ResourceConstraints
for _, r := range requestedResourceAccessIDs {
if rid := r.GetResourceID(); rid.Kind != resource.GetKind() || rid.Name != resource.GetName() {
continue
}
if c := r.GetConstraints(); c != nil {
constraints = c
break
}
}
var (
rolesForResource []types.Role
matchers []RoleMatcher
kubeResourceMatcher *KubeResourcesMatcher
)
kubernetesResources, err := getKubeResourcesFromResourceIDs(requestedResourceIDs, resource.GetName())
if err != nil {
return nil, trace.Wrap(err)
}
if len(kubernetesResources) > 0 {
kubeResourceMatcher = NewKubeResourcesMatcher(kubernetesResources)
matchers = append(matchers, kubeResourceMatcher)
}
switch rr := resource.(type) {
case types.Resource153UnwrapperT[IdentityCenterAccount]:
matchers = append(matchers, NewIdentityCenterAccountMatcher(rr.UnwrapT()))
case types.Resource153UnwrapperT[IdentityCenterAccountAssignment]:
matchers = append(matchers, NewIdentityCenterAccountAssignmentMatcher(rr.UnwrapT()))
}
// If ResourceConstraints were provided for this Resource, wrap existing
// matchers and add constraint-derived matchers. The wrapping gates
// principal-bearing matchers on the constraint's allowed set, while the
// constraint-derived matchers ensure roles are pruned to only those
// granting at least one of the constrained principals (e.g. SSH logins,
// AWS role ARNs).
if constraints != nil {
guard := WithConstraints(constraints)
for i := range matchers {
matchers[i] = guard(matchers[i])
}
constraintMatcher, err := MatcherFromConstraints(constraints)
if err != nil {
return nil, trace.Wrap(err)
}
if constraintMatcher != nil {
matchers = append(matchers, constraintMatcher)
}
}
for _, role := range allRoles {
roleAllowsAccess, err := m.roleAllowsResource(role, resource, loginHint, matchers...)
if err != nil {
return nil, trace.Wrap(err)
}
if !roleAllowsAccess {
// Role does not allow access to this resource. We will prune it
// unless it allows access to another resource.
continue
}
rolesForResource = append(rolesForResource, role)
}
// If any of the requested resources didn't match with the provided roles,
// we deny the request because the user is trying to request more access
// than what is allowed by its search_as_roles.
if kubeResourceMatcher != nil && len(kubeResourceMatcher.Unmatched()) > 0 {
resourcesStr, err := types.ResourceIDsToString(requestedResourceIDs)
if err != nil {
return nil, trace.Wrap(err)
}
return nil, trace.BadParameter(
`no roles configured in the "search_as_roles" for this user allow `+
`access to at least one requested resources. `+
`resources: %s roles: %v unmatched resources: %v`,
resourcesStr, roles, kubeResourceMatcher.Unmatched())
}
if len(loginHint) > 0 {
// If we have a login hint, request the single role with the fewest
// allowed logins. All roles at this point have already matched the
// requested login and will include it.
rolesForResource = fewestLogins(rolesForResource)
}
for _, role := range rolesForResource {
necessaryRoles[role.GetName()] = struct{}{}
}
}
if len(necessaryRoles) == 0 {
resourcesStr, err := types.ResourceIDsToString(requestedResourceIDs)
if err != nil {
return nil, trace.Wrap(err)
}
return nil, trace.BadParameter(
`no roles configured in the "search_as_roles" for this user allow `+
`access to any requested resources. The user may already have `+
`access to all requested resources with their existing roles. `+
`resources: %s roles: %v login: %q`,
resourcesStr, roles, loginHint)
}
prunedRoles := make([]string, 0, len(necessaryRoles))
for role := range necessaryRoles {
prunedRoles = append(prunedRoles, role)
}
return prunedRoles, nil
}
func fewestLogins(roles []types.Role) []types.Role {
if len(roles) == 0 {
return roles
}
fewest := roles[0]
fewestCount := countAllowedLogins(fewest)
for _, role := range roles[1:] {
if countAllowedLogins(role) < fewestCount {
fewest = role
}
}
return []types.Role{fewest}
}
func countAllowedLogins(role types.Role) int {
allowed := set.New(role.GetLogins(types.Allow)...)
for _, d := range role.GetLogins(types.Deny) {
allowed.Remove(d)
}
return allowed.Len()
}
func (m *RequestValidator) roleAllowsResource(
role types.Role,
resource types.ResourceWithLabels,
loginHint string,
extraMatchers ...RoleMatcher,
) (bool, error) {
roleSet := RoleSet{role}
var matchers []RoleMatcher
if len(loginHint) > 0 {
matchers = append(matchers, NewLoginMatcher(loginHint))
}
matchers = append(matchers, extraMatchers...)
_, err := roleSet.checkAccess(resource, m.userState.GetName(), m.userState.GetTraits(), AccessState{MFAVerified: true}, matchers...)
if trace.IsAccessDenied(err) {
// Access denied, this role does not allow access to this resource, no
// unexpected error to report.
return false, nil
}
if err != nil {
// Unexpected error, return it.
return false, trace.Wrap(err)
}
// Role allows access to this resource.
return true, nil
}
// getUnderlyingResourcesByResourceIDs gets the underlying resources the user
// requested access. Except for resource Kinds present in types.KubernetesResourcesKinds,
// the underlying resources are the same as requested. If the resource requested
// is a Kubernetes resource, we return the underlying Kubernetes cluster.
func (m *RequestValidator) getUnderlyingResourcesByResourceIDs(ctx context.Context, resourceIDs []types.ResourceID) ([]types.ResourceWithLabels, error) {
if len(resourceIDs) == 0 {
return []types.ResourceWithLabels{}, nil
}
// When searching for Kube Resources, we change the resource Kind to the Kubernetes
// Cluster in order to load the roles that grant access to it and to verify
// if the access to it is allowed. We later verify if every Kubernetes Resource
// requested is fulfilled by at least one role.
searchableResourcesIDs := slices.Clone(resourceIDs)
for i := range searchableResourcesIDs {
if slices.Contains(types.KubernetesResourcesKinds, searchableResourcesIDs[i].Kind) || strings.HasPrefix(searchableResourcesIDs[i].Kind, types.AccessRequestPrefixKindKube) {
searchableResourcesIDs[i].Kind = types.KindKubernetesCluster
}
}
// load the underlying resources.
resources, err := accessrequest.GetResourcesByResourceIDs(ctx, m.getter, searchableResourcesIDs)
return resources, trace.Wrap(err)
}
// getKubeResourcesFromResourceIDs returns the Kubernetes Resources requested for
// the configured cluster.
func getKubeResourcesFromResourceIDs(resourceIDs []types.ResourceID, clusterName string) ([]types.KubernetesResource, error) {
kubernetesResources := make([]types.KubernetesResource, 0, len(resourceIDs))
for _, resourceID := range resourceIDs {
if resourceID.Name != clusterName {
continue
}
// TODO(@creack): DELETE IN v20.0.0 when we no longer support legacy access request formats.
// Special case to support legacy "namespace" kind request.
if resourceID.Kind == types.KindKubeNamespace {
// If the target namespace is a wildcard, update the pattern to make sure cluster-wide resources won't be matched.
targetNS := resourceID.SubResourceName
if targetNS == types.Wildcard {
targetNS = "^.+$"
}
kubernetesResources = append(kubernetesResources,
types.KubernetesResource{
Kind: "namespaces",
Name: resourceID.SubResourceName,
APIGroup: "",
},
types.KubernetesResource{
Kind: types.Wildcard,
Name: types.Wildcard,
Namespace: targetNS,
APIGroup: "",
},
)
continue
}
if slices.Contains(types.KubernetesResourcesKinds, resourceID.Kind) || strings.HasPrefix(resourceID.Kind, types.AccessRequestPrefixKindKube) {
kind := types.KubernetesResourcesKindsPlurals[resourceID.Kind]
if kind == "" {
kind = resourceID.Kind
}
isClusterWide := slices.Contains(types.KubernetesClusterWideResourceKinds, resourceID.Kind) || strings.HasPrefix(kind, types.AccessRequestPrefixKindKubeClusterWide)
if !isClusterWide {
kind = strings.TrimPrefix(kind, types.AccessRequestPrefixKindKubeNamespaced)
} else {
kind = strings.TrimPrefix(kind, types.AccessRequestPrefixKindKubeClusterWide)
}
gk := schema.ParseGroupKind(kind)
if gk.Group == "" {
gk.Group = types.KubernetesResourcesV7KindGroups[resourceID.Kind]
}
switch {
case isClusterWide:
kubernetesResources = append(kubernetesResources, types.KubernetesResource{
Kind: gk.Kind,
Name: resourceID.SubResourceName,
APIGroup: gk.Group,
})
default:
splits := strings.Split(resourceID.SubResourceName, "/")
if len(splits) != 2 {
return nil, trace.BadParameter("subresource name %q does not follow <namespace>/<name> format", resourceID.SubResourceName)
}
kubernetesResources = append(kubernetesResources, types.KubernetesResource{
Kind: gk.Kind,
Namespace: splits[0],
Name: splits[1],
APIGroup: gk.Group,
})
}
}
}
return kubernetesResources, nil
}
func newReviewPermissionParser() (*typical.Parser[reviewPermissionContext, bool], error) {
return typical.NewParser[reviewPermissionContext, bool](typical.ParserSpec[reviewPermissionContext]{
Variables: map[string]typical.Variable{
"reviewer.roles": typical.DynamicVariable(func(ctx reviewPermissionContext) ([]string, error) {
return ctx.reviewer.roles, nil
}),
"reviewer.traits": typical.DynamicVariable(func(ctx reviewPermissionContext) (map[string][]string, error) {
return ctx.reviewer.traits, nil
}),
"request.roles": typical.DynamicVariable(func(ctx reviewPermissionContext) ([]string, error) {
return ctx.request.roles, nil
}),
"request.reason": typical.DynamicVariable(func(ctx reviewPermissionContext) (string, error) {
return ctx.request.reason, nil
}),
"request.system_annotations": typical.DynamicVariable(func(ctx reviewPermissionContext) (map[string][]string, error) {
return ctx.request.systemAnnotations, nil
}),
},
Functions: map[string]typical.Function{
"equals": typical.BinaryFunction[reviewPermissionContext](equalsFunc),
"contains": typical.BinaryFunction[reviewPermissionContext](containsFunc),
"regexp.match": typical.BinaryFunction[reviewPermissionContext](regexpMatchFunc),
},
})
}
func mustNewReviewPermissionParser() *typical.Parser[reviewPermissionContext, bool] {
parser, err := newReviewPermissionParser()
if err != nil {
panic(err)
}
return parser
}
var (
reviewPermissionParser = mustNewReviewPermissionParser()
)
func parseReviewPermissionExpression(expr string) (typical.Expression[reviewPermissionContext, bool], error) {
parsed, err := reviewPermissionParser.Parse(expr)
return parsed, trace.Wrap(err, "parsing review.where expression")
}
func newThresholdFilterParser() (*typical.Parser[thresholdFilterContext, bool], error) {
return typical.NewParser[thresholdFilterContext, bool](typical.ParserSpec[thresholdFilterContext]{
Variables: map[string]typical.Variable{
"reviewer.roles": typical.DynamicVariable(func(ctx thresholdFilterContext) ([]string, error) {
return ctx.reviewer.roles, nil
}),
"reviewer.traits": typical.DynamicVariable(func(ctx thresholdFilterContext) (map[string][]string, error) {
return ctx.reviewer.traits, nil
}),
"review.reason": typical.DynamicVariable(func(ctx thresholdFilterContext) (string, error) {
return ctx.review.reason, nil
}),
"review.annotations": typical.DynamicVariable(func(ctx thresholdFilterContext) (map[string][]string, error) {
return ctx.review.annotations, nil
}),
"request.roles": typical.DynamicVariable(func(ctx thresholdFilterContext) ([]string, error) {
return ctx.request.roles, nil
}),
"request.reason": typical.DynamicVariable(func(ctx thresholdFilterContext) (string, error) {
return ctx.request.reason, nil
}),
"request.system_annotations": typical.DynamicVariable(func(ctx thresholdFilterContext) (map[string][]string, error) {
return ctx.request.systemAnnotations, nil
}),
},
Functions: map[string]typical.Function{
"equals": typical.BinaryFunction[thresholdFilterContext](equalsFunc),
"contains": typical.BinaryFunction[thresholdFilterContext](containsFunc),
"regexp.match": typical.BinaryFunction[thresholdFilterContext](regexpMatchFunc),
},
})
}
func mustNewThresholdFilterParser() *typical.Parser[thresholdFilterContext, bool] {
parser, err := newThresholdFilterParser()
if err != nil {
panic(err)
}
return parser
}
var (
thresholdFilterParser = mustNewThresholdFilterParser()
)
func parseThresholdFilterExpression(expr string) (typical.Expression[thresholdFilterContext, bool], error) {
parsed, err := thresholdFilterParser.Parse(expr)
return parsed, trace.Wrap(err, "parsing threshold filter expression")
}
func equalsFunc(a, b any) (bool, error) {
switch aval := a.(type) {
case string:
bval, ok := b.(string)
if ok {
return aval == bval, nil
}
case []string:
bval, ok := b.([]string)
if ok {
return slices.Equal(aval, bval), nil
}
}
return false, trace.BadParameter("parameter types must match and be string or []string, got (%T, %T)", a, b)
}
func containsFunc(s []string, v string) (bool, error) {
return slices.Contains(s, v), nil
}
func regexpMatchFunc(list []string, re string) (bool, error) {
match, err := utils.RegexMatchesAny(list, re)
if err != nil {
return false, trace.Wrap(err, "invalid regular expression %q", re)
}
return match, nil
}
/*
* Teleport
* Copyright (C) 2024 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"
"fmt"
"log/slog"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/client/proto"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/sortcache"
)
type accessRequestCacheIndex string
const (
// accessRequestID is the name of the sort index used for sorting request by ID (equivalent to proto.AccessRequestSort_DEFAULT since
// access requests currently default to being sorted by ID in the backend).
accessRequestID accessRequestCacheIndex = "ID"
// accessRequestCreated is the name of the sort index used for sorting requests by creation time (this is typically the sort order
// used in user interfaces, since most users that want to view requests want to see the most recent requests specifically).
accessRequestCreated accessRequestCacheIndex = "Created"
// accessRequestState is the name of the sort index used for sorting requests by their current state (pending, approved, etc).
accessRequestState accessRequestCacheIndex = "State"
// accessRequestUser is the name of the sort index used for sorting requests by the person who created the request.
accessRequestUser accessRequestCacheIndex = "User"
)
// AccessRequestCacheConfig holds the configuration parameters for an [AccessRequestCache].
type AccessRequestCacheConfig struct {
// Clock is a clock for time-related operation.
Clock clockwork.Clock
// Events is an event system client.
Events types.Events
// Getter is an access request getter client.
Getter AccessRequestGetter
// MaxRetryPeriod is the maximum retry period on failed watches.
MaxRetryPeriod time.Duration
}
// CheckAndSetDefaults valides the config and provides reasonable defaults for optional fields.
func (c *AccessRequestCacheConfig) CheckAndSetDefaults() error {
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
if c.Events == nil {
return trace.BadParameter("access request cache config missing event system client")
}
if c.Getter == nil {
return trace.BadParameter("access request cache config missing access request getter")
}
return nil
}
// AccessRequestCache is a custom cache for access requests that offers custom sort indexes not
// supported by the standard backend implementation. As with all caches, the state observed during
// reads may be slightly outdated. There is a builtin fallback that always routes requests for a
// single specific access request (specified by ID) to the real backend, to avoid outdated single-resource
// reads. Usecases that need perfectly up to date information (e.g. loading an access request in order
// to generate a certificate) should always load the desired request by ID for this reason.
type AccessRequestCache struct {
rw sync.RWMutex
cfg AccessRequestCacheConfig
primaryCache *sortcache.SortCache[*types.AccessRequestV3, accessRequestCacheIndex]
ttlCache *utils.FnCache
initC chan struct{}
initOnce sync.Once
closeContext context.Context
cancel context.CancelFunc
// onInit is a callback used in tests to detect
// individual initializations.
onInit func()
}
// NewAccessRequestCache sets up a new [AccessRequestCache] instance based on the supplied
// configuration. The cache is initialized asychronously in the background, so while it is
// safe to read from it immediately, performance is better after the cache properly initializes.
func NewAccessRequestCache(cfg AccessRequestCacheConfig) (*AccessRequestCache, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
ctx, cancel := context.WithCancel(context.Background())
ttlCache, err := utils.NewFnCache(utils.FnCacheConfig{
Context: ctx,
TTL: 15 * time.Second,
Clock: cfg.Clock,
})
if err != nil {
cancel()
return nil, trace.Wrap(err)
}
c := &AccessRequestCache{
cfg: cfg,
ttlCache: ttlCache,
initC: make(chan struct{}),
closeContext: ctx,
cancel: cancel,
}
if _, err := newResourceWatcher(ctx, c, ResourceWatcherConfig{
Component: "access-request-cache",
Client: cfg.Events,
MaxRetryPeriod: cfg.MaxRetryPeriod,
}); err != nil {
cancel()
return nil, trace.Wrap(err)
}
return c, nil
}
// ListAccessRequests is an access request getter with pagination and sorting options.
func (c *AccessRequestCache) ListAccessRequests(ctx context.Context, req *proto.ListAccessRequestsRequest) (*proto.ListAccessRequestsResponse, error) {
rsp, err := c.ListMatchingAccessRequests(ctx, req, func(_ *types.AccessRequestV3) bool {
return true
})
return rsp, trace.Wrap(err)
}
// ListMatchingAccessRequests is equivalent to ListAccessRequests except that it adds the ability to provide an arbitrary matcher function. This method
// should be preferred when using custom filtering (e.g. access-controls), since the paginations keys used by the access request cache are non-standard.
func (c *AccessRequestCache) ListMatchingAccessRequests(ctx context.Context, req *proto.ListAccessRequestsRequest, match func(*types.AccessRequestV3) bool) (*proto.ListAccessRequestsResponse, error) {
const maxPageSize = 16_000
if req.Filter == nil {
req.Filter = &types.AccessRequestFilter{}
}
if req.Filter.ID != "" {
// important special case: single-request lookups must always be forwarded to the real backend to avoid race conditions whereby
// stale cache state causes spurious errors due to users trying to utilize an access request immediately after it gets approved.
rsp, err := c.cfg.Getter.ListAccessRequests(ctx, req)
if err != nil {
return nil, trace.Wrap(err)
}
// fallback doesn't apply the match function, so we need to apply it manually here.
matched := rsp.AccessRequests[:0]
for _, req := range rsp.AccessRequests {
if !match(req) {
continue
}
matched = append(matched, req)
}
rsp.AccessRequests = matched
return rsp, nil
}
if req.Limit == 0 {
req.Limit = apidefaults.DefaultChunkSize
}
if req.Limit > maxPageSize {
return nil, trace.BadParameter("page size of %d is too large", req.Limit)
}
cache, err := c.read(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
var index accessRequestCacheIndex
switch req.Sort {
case proto.AccessRequestSort_DEFAULT:
index = accessRequestID
case proto.AccessRequestSort_CREATED:
index = accessRequestCreated
case proto.AccessRequestSort_STATE:
index = accessRequestState
case proto.AccessRequestSort_USER:
index = accessRequestUser
default:
return nil, trace.BadParameter("unsupported access request sort index '%v'", req.Sort)
}
if !cache.HasIndex(index) {
// this case would be a fairly trivial programming error, but its best to give it
// a friendly error message.
return nil, trace.Errorf("access request cache was not configured with sort index %q (this is a bug)", index)
}
accessRequests := cache.Ascend
if req.Descending {
accessRequests = cache.Descend
}
limit := int(req.Limit)
// perform the traversal until we've seen all items or fill the page
var rsp proto.ListAccessRequestsResponse
now := time.Now()
var expired int
for r := range accessRequests(index, req.StartKey, "") {
if len(rsp.AccessRequests) == limit {
rsp.NextKey = cache.KeyOf(index, r)
break
}
if !r.Expiry().IsZero() && now.After(r.Expiry()) {
expired++
// skip requests that appear expired. some backends can take up to 48 hours to expired items
// and access requests showing up past their expiry time is particularly confusing.
continue
}
if !req.Filter.Match(r) || !match(r) {
continue
}
c := r.Copy()
cr, ok := c.(*types.AccessRequestV3)
if !ok {
slog.WarnContext(ctx, "clone returned unexpected type (this is a bug)", "expected", logutils.TypeAttr(r), "got", logutils.TypeAttr(c))
continue
}
rsp.AccessRequests = append(rsp.AccessRequests, cr)
}
if expired > 0 {
// this is a debug-level log since some amount of delay between expiry and backend cleanup is expected, but
// very large and/or disproportionate numbers of stale access requests might be a symptom of a deeper issue.
slog.DebugContext(ctx, "omitting expired access requests from cache read", "count", expired)
}
return &rsp, nil
}
// fetch configures a sortcache and inserts all currently extant access requests into it. this method is used both
// as the means of setting up the initial primary cache state, and for creating temporary cache states to read from
// when the primary is unhealthy.
func (c *AccessRequestCache) fetch(ctx context.Context) (*sortcache.SortCache[*types.AccessRequestV3, accessRequestCacheIndex], error) {
cache := sortcache.New(sortcache.Config[*types.AccessRequestV3, accessRequestCacheIndex]{
Indexes: map[accessRequestCacheIndex]func(*types.AccessRequestV3) string{
accessRequestID: func(req *types.AccessRequestV3) string {
// since accessRequestID is equivalent to the DEFAULT sort index (i.e. the sort index of the backend),
// it is preferable to keep its format equivalent to the format of the NextKey/StartKey values
// expected by the backend implementation of ListResources.
return req.GetName()
},
accessRequestCreated: func(req *types.AccessRequestV3) string {
return fmt.Sprintf("%s/%s", req.GetCreationTime().Format(time.RFC3339), req.GetName())
},
accessRequestState: func(req *types.AccessRequestV3) string {
return fmt.Sprintf("%s/%s", req.GetState().String(), req.GetName())
},
accessRequestUser: func(req *types.AccessRequestV3) string {
return fmt.Sprintf("%s/%s", req.GetUser(), req.GetName())
},
},
})
var req proto.ListAccessRequestsRequest
for {
rsp, err := c.cfg.Getter.ListAccessRequests(ctx, &req)
if err != nil {
return nil, trace.Wrap(err)
}
for _, r := range rsp.AccessRequests {
if evicted := cache.Put(r); evicted != 0 {
// this warning, if it appears, means that we configured our indexes incorrectly and one access request is overwriting another.
// the most likely explanation is that one of our indexes is missing the request id suffix we typically use.
slog.WarnContext(ctx, "conflict during access request fetch (this is a bug and may result in missing requests)", "id", r.GetName(), "evicted", evicted)
}
}
if rsp.NextKey == "" {
break
}
req.StartKey = rsp.NextKey
}
return cache, nil
}
// read gets a read-only view into a valid cache state. it prefers reading from the primary cache, but will fallback
// to a periodically reloaded temporary state when the primary state is unhealthy.
func (c *AccessRequestCache) read(ctx context.Context) (*sortcache.SortCache[*types.AccessRequestV3, accessRequestCacheIndex], error) {
c.rw.RLock()
primary := c.primaryCache
c.rw.RUnlock()
// primary cache state is healthy, so use that. note that we don't protect access to the sortcache itself
// via our rw lock. sortcaches have their own internal locking. we just use our lock to protect the *pointer*
// to the sortcache.
if primary != nil {
return primary, nil
}
temp, err := utils.FnCacheGet(ctx, c.ttlCache, "access-request-cache", func(ctx context.Context) (*sortcache.SortCache[*types.AccessRequestV3, accessRequestCacheIndex], error) {
return c.fetch(ctx)
})
// primary may have been concurrently loaded. if it was, prefer using that.
c.rw.RLock()
primary = c.primaryCache
c.rw.RUnlock()
if primary != nil {
return primary, nil
}
return temp, trace.Wrap(err)
}
// --- the below methods implement the resourceCollector interface ---
// resourceKinds is part of the resourceCollector interface and is used to configure the event watcher
// that monitors for access request modifications.
func (c *AccessRequestCache) resourceKinds() []types.WatchKind {
return []types.WatchKind{
{
Kind: types.KindAccessRequest,
},
}
}
// getResourcesAndUpdateCurrent is part of the resourceCollector interface and is called one the
// event stream for the cache has been initialized to trigger setup of the initial primary cache state.
func (c *AccessRequestCache) getResourcesAndUpdateCurrent(ctx context.Context) error {
cache, err := c.fetch(ctx)
if err != nil {
return trace.Wrap(err)
}
c.rw.Lock()
defer c.rw.Unlock()
c.primaryCache = cache
c.initOnce.Do(func() {
close(c.initC)
})
if c.onInit != nil {
c.onInit()
}
return nil
}
// SetInitCallback is used in tests that care about cache inits.
func (c *AccessRequestCache) SetInitCallback(cb func()) {
c.rw.Lock()
defer c.rw.Unlock()
c.onInit = cb
}
// processEventsAndUpdateCurrent is part of the resourceCollector interface and is used to update the
// primary cache state when modification events occur.
func (c *AccessRequestCache) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
if len(events) < 1 {
return
}
c.rw.RLock()
cache := c.primaryCache
c.rw.RUnlock()
if cache == nil {
return
}
for _, event := range events {
switch event.Type {
case types.OpPut:
req, ok := event.Resource.(*types.AccessRequestV3)
if !ok {
slog.WarnContext(ctx, "unexpected resource type in event", "expected", logutils.TypeAttr(req), "got", logutils.TypeAttr(event.Resource))
continue
}
if evicted := cache.Put(req); evicted > 1 {
// this warning, if it appears, means that we configured our indexes incorrectly and one access request is overwriting another.
// the most likely explanation is that one of our indexes is missing the request id suffix we typically use.
slog.WarnContext(ctx, "request put event resulted in multiple cache evictions (this is a bug)", "id", req.GetName(), "evicted", evicted)
}
case types.OpDelete:
cache.Delete(accessRequestID, event.Resource.GetName())
default:
slog.WarnContext(ctx, "unexpected event variant", "op", logutils.StringerAttr(event.Type), "resource", logutils.TypeAttr(event.Resource))
}
}
}
// notifyStale is part of the resourceCollector interface and is used to inform
// the access request cache that its view is outdated (presumably due to issues with
// the event stream).
func (c *AccessRequestCache) notifyStale() {
c.rw.Lock()
defer c.rw.Unlock()
if c.primaryCache == nil {
return
}
c.primaryCache = nil
c.initC = make(chan struct{})
c.initOnce = sync.Once{}
}
// initializationChan is part of the resourceCollector interface and gets the channel
// used to signal that the accessRequestCache has been initialized.
func (c *AccessRequestCache) initializationChan() <-chan struct{} {
c.rw.RLock()
defer c.rw.RUnlock()
return c.initC
}
// InitializationChan is part of the resourceCollector interface and gets the channel
// used to signal that the accessRequestCache has been initialized.
func (c *AccessRequestCache) InitializationChan() <-chan struct{} {
return c.initializationChan()
}
// Close terminates the background process that keeps the access request cache up to
// date, and terminates any inflight load operations.
func (c *AccessRequestCache) Close() error {
c.cancel()
return nil
}
/*
* 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 (
"cmp"
"context"
"crypto/x509"
"fmt"
"iter"
"log/slog"
"net"
"net/http"
"net/url"
"os"
"slices"
"strconv"
"strings"
"sync"
"github.com/gravitational/trace"
"github.com/spiffe/go-spiffe/v2/spiffeid"
"golang.org/x/net/idna"
corev1 "k8s.io/api/core/v1"
"k8s.io/apimachinery/pkg/util/validation"
kyaml "k8s.io/apimachinery/pkg/util/yaml"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/wrappers"
"github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/api/utils/tlsutils"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/scopes"
scopedapp "github.com/gravitational/teleport/lib/scopes/app"
"github.com/gravitational/teleport/lib/utils"
)
// AppGetter defines interface for fetching application resources.
type AppGetter interface {
// GetApps returns all application resources.
GetApps(context.Context) ([]types.Application, error)
// ListApps returns a page of application resources.
ListApps(ctx context.Context, limit int, startKey string) ([]types.Application, string, error)
// Apps returns application resources within the range [start, end).
Apps(ctx context.Context, start, end string) iter.Seq2[types.Application, error]
// GetApp returns the specified application resource.
GetApp(ctx context.Context, name string) (types.Application, error)
}
// Applications defines an interface for managing application resources.
type Applications interface {
// AppGetter provides methods for fetching application resources.
AppGetter
// CreateApp creates a new application resource.
CreateApp(context.Context, types.Application) error
// UpdateApp updates an existing application resource.
UpdateApp(context.Context, types.Application) error
// DeleteApp removes the specified application resource.
DeleteApp(ctx context.Context, name string) error
// DeleteAllApps removes all database resources.
DeleteAllApps(context.Context) error
}
// ApplicationsInternal extends the Access interface with auth-specific internal methods.
type ApplicationsInternal interface {
Applications
// AppendPutAppActions adds conditional actions to an atomic write to create
// or update an application resource.
AppendPutAppActions(
actions []backend.ConditionalAction,
app types.Application,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteAppActions adds conditional actions to an atomic write to
// delete an application resource.
AppendDeleteAppActions(
actions []backend.ConditionalAction,
name string,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
// TODO(williamo/scopes): remove once dynamic scoped app registration is
// supported.
func EnsureNotScopedApp(app types.Application) error {
if app == nil {
return trace.BadParameter("nil application")
}
if scope := app.GetScope(); scope != "" {
return trace.BadParameter("application %q cannot be created with scope %q: dynamic registration of scoped applications is not supported, remove the scope attribute", app.GetName(), scope)
}
return nil
}
// ValidateApp checks an Application's name, public_addr, and
// required_apps.
func ValidateApp(app types.Application, proxyGetter ProxyGetter) error {
if app == nil {
return trace.BadParameter("nil application")
}
// Subdomain (not label) so integrations can produce dotted names.
// Allow underscores: Cloud and self-hosted both already accept
// underscored app names, and curl/Go-based clients route them via
// the wildcard cert. Stricter rejection would break callers like
// the Terraform provider whose test fixtures use snake_case.
if errs := validation.IsDNS1123SubdomainWithUnderscore(app.GetName()); len(errs) > 0 {
return trace.BadParameter("application name %q must be a valid DNS name (lowercase alphanumeric, '-', '_', or '.', must start and end with alphanumeric, max 253 chars): https://goteleport.com/docs/enroll-resources/application-access/guides/connecting-apps/#application-name", app.GetName())
}
// required_apps lookup is exact-string; a mixed-case entry would
// never match a lowercased primary name.
for _, required := range app.GetRequiredAppNames() {
if errs := validation.IsDNS1123SubdomainWithUnderscore(required); len(errs) > 0 {
return trace.BadParameter("application %q references required_apps entry %q which must be a valid DNS name (lowercase alphanumeric, '-', '_', or '.', start and end alphanumeric, max 253 chars): https://goteleport.com/docs/enroll-resources/application-access/guides/connecting-apps/#application-name", app.GetName(), required)
}
}
if app.GetTLS() != nil {
if err := validateAppTLS(app); err != nil {
return trace.Wrap(err)
}
}
if region := app.GetAWSRegion(); region != "" {
if err := aws.IsValidRegion(region); err != nil {
return trace.BadParameter(
"Application %q is configured with an invalid AWS region (%q)",
app.GetName(),
region,
)
}
}
if scope := app.GetScope(); scope != "" {
if !scopedapp.ScopedAppPublicAddrValid(scope, app.GetName(), app.GetPublicAddr()) {
return trace.BadParameter("scoped app %q public address %q does not match its derived address for scope %q", app.GetName(), app.GetPublicAddr(), scope)
}
}
if app.GetPublicAddr() == "" {
return nil
}
if err := ValidatePublicAddr(app.GetName(), app.GetPublicAddr()); err != nil {
return trace.Wrap(err)
}
appAddr, err := utils.ParseAddr(app.GetPublicAddr())
if err != nil {
return trace.Wrap(err)
}
// Normalize to ASCII for the proxy-collision compare below; the
// proxy public_addr is not run through ValidatePublicAddr.
asciiAppHostname, err := idna.ToASCII(strings.TrimRight(appAddr.Host(), "."))
if err != nil {
return trace.Wrap(err, "app %q has an invalid IDN hostname %q", app.GetName(), appAddr.Host())
}
proxyServers, err := clientutils.CollectWithFallback(context.TODO(), proxyGetter.ListProxyServers, func(context.Context) ([]types.Server, error) {
//nolint:staticcheck // TODO(kiosion) DELETE IN 21.0.0
return proxyGetter.GetProxies()
})
if err != nil {
return trace.Wrap(err)
}
// Prevent routing conflicts and session hijacking by ensuring the application's public address does not match the
// public address of any proxy. If an application shares a public address with a proxy, requests intended for the
// proxy could be misrouted to the application, compromising security.
for _, proxyServer := range proxyServers {
proxyAddrs, err := utils.ParseAddrs(proxyServer.GetPublicAddrs())
if err != nil {
return trace.Wrap(err)
}
for _, proxyAddr := range proxyAddrs {
// Also convert the proxy's public address hostname to its ASCII representation for comparison and strip any
// trailing dots.
asciiProxyHostname, err := idna.ToASCII(strings.TrimRight(proxyAddr.Host(), "."))
if err != nil {
return trace.Wrap(err, "proxy %q has an invalid IDN hostname %q", proxyServer.GetName(), proxyAddr)
}
// Compare the ASCII-normalized hostnames for equality, ignoring case.
if strings.EqualFold(asciiProxyHostname, asciiAppHostname) {
return trace.BadParameter(
"Application %q public address %q conflicts with the Teleport Proxy public address. "+
"Configure the application to use a unique public address that does not match the proxy's public addresses. "+
"Refer to https://goteleport.com/docs/enroll-resources/application-access/guides/connecting-apps/#customize-public-address.",
app.GetName(),
app.GetPublicAddr(),
)
}
}
}
return nil
}
// validateAppTLS validates application TLS options.
func validateAppTLS(a types.Application) error {
if !types.AppSupportsTLSConfig(a.GetURI()) {
return trace.BadParameter(
"App %q can only specify 'tls' settings for URI schemes that use upstream TLS. Supported schemes are: %s",
a.GetName(),
quoteAndJoin(types.AppSchemesWithTLSSupport),
)
}
tls := a.GetTLS()
var mode types.AppTLSMode
switch tls.Mode {
case types.AppTLSModeInsecure,
types.AppTLSModeVerifyFull,
types.AppTLSModeVerifyServerName,
types.AppTLSModeVerifySpiffeID:
mode = tls.Mode
case "":
// When not specified, use the evaluated mode.
mode = a.GetTLSMode()
default:
return trace.BadParameter(
"App %q has invalid 'tls.mode' %q. Supported values are: %s",
a.GetName(),
tls.Mode,
quoteAndJoin([]string{
types.AppTLSModeInsecure,
types.AppTLSModeVerifyFull,
types.AppTLSModeVerifyServerName,
types.AppTLSModeVerifySpiffeID,
}),
)
}
if a.GetInsecureSkipVerify() && mode != types.AppTLSModeInsecure {
return trace.BadParameter(
"App %q cannot specify 'insecure_skip_verify: true' (deprecated) and 'tls.mode: %q'. Drop 'insecure_skip_verify', and if you want the app to use insecure connections set 'tls.mode: %q'",
a.GetName(),
mode,
types.AppTLSModeInsecure,
)
}
switch tls.ClientCertMode {
case types.AppClientCertModeManaged:
if mode == types.AppTLSModeInsecure {
return trace.BadParameter("App %q can only enable 'tls.client_cert_mode' when 'tls.mode' is %q", a.GetName(), types.AppTLSModeVerifyFull)
}
case types.AppClientCertModeDisabled, "":
default:
return trace.BadParameter(
"App %q has invalid 'tls.client_cert_mode'. Supported values are: %s",
a.GetName(),
quoteAndJoin([]string{"", types.AppClientCertModeDisabled, types.AppClientCertModeManaged}),
)
}
switch mode {
case types.AppTLSModeVerifyFull:
// Note: tls.ServerName is optional and doesn't require any specific validation.
if err := isValidSpiffeID(tls.ServerSpiffeId); err != nil {
return trace.BadParameter("App %q has invalid `tls.server_spiffe_id`. The SPIFFE ID must be complete (trust domain and path) and start with 'spiffe://': %v", a.GetName(), err)
}
case types.AppTLSModeVerifyServerName:
// Note: tls.ServerName is optional and doesn't require any specific validation.
if tls.ServerSpiffeId != "" {
return trace.BadParameter("App %q 'tls.server_spiffe_id' is not used when mode is set to %q. To perform both, server name and SPIFFE ID verifications use %q mode", a.GetName(), mode, types.AppTLSModeVerifyFull)
}
case types.AppTLSModeVerifySpiffeID:
if err := isValidSpiffeID(tls.ServerSpiffeId); err != nil {
return trace.BadParameter("App %q has invalid `tls.server_spiffe_id`. The SPIFFE ID must be complete (trust domain and path) and start with 'spiffe://': %v", a.GetName(), err)
}
if tls.ServerName != "" {
return trace.BadParameter("App %q 'tls.server_name' is not used when mode is set to %q. To perform both, server name and SPIFFE ID verifications use %q mode", a.GetName(), mode, types.AppTLSModeVerifyFull)
}
case types.AppTLSModeInsecure:
if tls.ServerName != "" || tls.ServerSpiffeId != "" || len(tls.AllowedCas) > 0 {
return trace.BadParameter("App %q 'tls' are not in use since mode is set to %q", a.GetName(), mode)
}
}
supportedCAs := types.AppSupportedInternalCAs()
for _, allowedCA := range tls.AllowedCas {
if slices.Contains(supportedCAs, allowedCA) {
continue
}
if err := isValidCACertificatePEM(allowedCA); err != nil {
return trace.BadParameter(
"App %q 'tls.allowed_cas' values must include valid PEM-encoded CA certificates or a Teleport CA alias (%s): %s",
a.GetName(),
quoteAndJoin(supportedCAs),
err,
)
}
}
return nil
}
// ValidateAppServer checks the outer AppServer name (the backend
// storage key) and delegates to ValidateApp for the inner app. The
// outer name can differ from the inner name when an admin creates
// an app_server resource directly (for example, via `tctl create`),
// so it needs its own DNS-1123 check.
func ValidateAppServer(server types.AppServer, proxyGetter ProxyGetter) error {
if server == nil {
return trace.BadParameter("nil app server")
}
if errs := validation.IsDNS1123SubdomainWithUnderscore(server.GetName()); len(errs) > 0 {
return trace.BadParameter("app server name %q must be a valid DNS name (lowercase alphanumeric, '-', '_', or '.', must start and end with alphanumeric, max 253 chars): %s", server.GetName(), strings.Join(errs, ", "))
}
app := server.GetApp()
if app != nil && !AppServerScopesEqual(server.GetScope(), app.GetScope()) {
return trace.BadParameter("app server %q scope %q does not match its embedded app scope %q", server.GetName(), server.GetScope(), app.GetScope())
}
return trace.Wrap(ValidateApp(app, proxyGetter))
}
// AppServerScopesEqual reports whether an app server's scope and its embedded
// app's scope are equivalent.
func AppServerScopesEqual(serverScope, appScope string) bool {
// Empty string comparison is treated as orthogonal in scopes.Compare.
// If server scope is empty (unscoped), we should make sure the app's scope is also empty, and vice versa.
if serverScope == "" || appScope == "" {
return serverScope == appScope
}
return scopes.Compare(serverScope, appScope) == scopes.Equivalent
}
// GetCursorForAppServer returns the resource cursor identifying an app server
// in the logical resource stream: "<host-id>/<name>" for unscoped app servers
// and "~scoped/<encoded-scope>/<host-id>/<name>" for scoped apps.
func GetCursorForAppServer(server types.AppServer) string {
return scopes.MakeResourceCursorWithHost(server.GetScope(), server.GetHostID(), server.GetName())
}
// ValidatePublicAddr requires a lowercase DNS-1123 hostname. An
// empty addr is treated as unset.
func ValidatePublicAddr(appName, addr string) error {
if addr == "" {
return nil
}
// IPv4 literals satisfy IsDNS1123Subdomain (digits + dots);
// reject explicitly so routing-by-hostname holds.
if net.ParseIP(addr) != nil {
return trace.BadParameter("application %q public_addr %q must not be an IP address, Teleport Application Access uses DNS names for routing", appName, addr)
}
if errs := validation.IsDNS1123Subdomain(addr); len(errs) > 0 {
return trace.BadParameter("application %q public_addr %q must be a valid DNS name (lowercase alphanumeric, '-', or '.', no trailing dot, no IDN Unicode -- use punycode): %s", appName, addr, strings.Join(errs, ", "))
}
return nil
}
// NormalizeAppServerForHeartbeat case-folds a legacy agent's name and
// strips a scheme/port from its public_addr. Heartbeat-only; admin
// paths must not call it.
func NormalizeAppServerForHeartbeat(server types.AppServer) {
app := server.GetApp()
if app == nil {
return
}
appName := strings.ToLower(app.GetName())
if appName != app.GetName() {
app.SetName(appName)
}
serverName := server.GetName()
if strings.EqualFold(serverName, appName) && serverName != appName {
server.SetName(appName)
}
normalizedPublicAddr := normalizeHeartbeatPublicAddr(app.GetPublicAddr())
if normalizedPublicAddr != app.GetPublicAddr() {
app.SetPublicAddr(normalizedPublicAddr)
}
// required_apps go through the same DNS-1123 check in ValidateApp;
// lowercase legacy mixed-case entries so older agents keep working.
if appV3, ok := app.(*types.AppV3); ok {
for i, required := range appV3.Spec.RequiredAppNames {
appV3.Spec.RequiredAppNames[i] = strings.ToLower(required)
}
}
}
// normalizeHeartbeatPublicAddr strips a URL scheme, path, or port
// from a legacy public_addr and lowercases the result.
func normalizeHeartbeatPublicAddr(addr string) string {
if addr == "" {
return addr
}
if strings.Contains(addr, "://") {
if u, err := url.Parse(addr); err == nil && u.Hostname() != "" {
return strings.ToLower(u.Hostname())
}
return addr
}
if host, _, err := net.SplitHostPort(addr); err == nil {
return strings.ToLower(host)
}
return strings.ToLower(addr)
}
// MarshalApp marshals Application resource to JSON.
func MarshalApp(app types.Application, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch app := app.(type) {
case *types.AppV3:
if err := app.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, app))
default:
return nil, trace.BadParameter("unsupported app resource %T", app)
}
}
// UnmarshalApp unmarshals Application resource from JSON.
func UnmarshalApp(data []byte, opts ...MarshalOption) (types.Application, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing app resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var app types.AppV3
if err := utils.FastUnmarshal(data, &app); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := app.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
app.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
app.SetExpiry(cfg.Expires)
}
return &app, nil
}
return nil, trace.BadParameter("unsupported app resource version %q", h.Version)
}
// MarshalAppServer marshals the AppServer resource to JSON.
func MarshalAppServer(appServer types.AppServer, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch appServer := appServer.(type) {
case *types.AppServerV3:
if err := appServer.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, appServer))
default:
return nil, trace.BadParameter("unsupported app server resource %T", appServer)
}
}
// UnmarshalAppServer unmarshals AppServer resource from JSON.
func UnmarshalAppServer(data []byte, opts ...MarshalOption) (types.AppServer, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing app server data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var s types.AppServerV3
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("unsupported app server resource version %q", h.Version)
}
// NewApplicationFromKubeService creates application resources from kubernetes service.
// It transforms service fields and annotations into appropriate Teleport app fields.
// Service labels are copied to app labels.
func NewApplicationFromKubeService(service corev1.Service, clusterName, protocol string, port corev1.ServicePort) (types.Application, error) {
appURI := buildAppURI(protocol, GetServiceFQDN(service), service.GetAnnotations()[types.DiscoveryPathLabel], port.Port)
rewriteConfig, err := getAppRewriteConfig(service.GetAnnotations())
if err != nil {
return nil, trace.Wrap(err, "could not get app rewrite config for the service")
}
appNameAnnotation := service.GetAnnotations()[types.DiscoveryAppNameLabel]
appName, err := getAppName(service.GetName(), service.GetNamespace(), clusterName, port.Name, appNameAnnotation)
if err != nil {
return nil, trace.Wrap(err, "could not create app name for the service")
}
labels, err := getAppLabels(service.GetLabels(), clusterName)
if err != nil {
return nil, trace.Wrap(err, "could not get labels for the service")
}
app, err := types.NewAppV3(types.Metadata{
Name: appName,
Description: cmp.Or(
getDescription(service.GetAnnotations()),
fmt.Sprintf("Discovered application in Kubernetes cluster %q", clusterName),
),
Labels: labels,
}, types.AppSpecV3{
URI: appURI,
Rewrite: rewriteConfig,
InsecureSkipVerify: getTLSInsecureSkipVerify(service.GetAnnotations()),
PublicAddr: getPublicAddr(service.GetAnnotations()),
})
if err != nil {
return nil, trace.Wrap(err, "could not create an app from Kubernetes service")
}
return app, nil
}
// GetServiceFQDN returns the fully qualified domain name for the service.
func GetServiceFQDN(service corev1.Service) string {
// If service type is ExternalName it points to external DNS name, to keep correct
// HOST for HTTP requests we return already final external DNS name.
// https://kubernetes.io/docs/concepts/services-networking/service/#externalname
if service.Spec.Type == corev1.ServiceTypeExternalName {
return service.Spec.ExternalName
}
return fmt.Sprintf("%s.%s.svc.%s", service.GetName(), service.GetNamespace(), clusterDomainResolver())
}
func buildAppURI(protocol, serviceFQDN, path string, port int32) string {
return (&url.URL{
Scheme: protocol,
Host: net.JoinHostPort(serviceFQDN, strconv.Itoa(int(port))),
Path: path,
}).String()
}
func getAppRewriteConfig(annotations map[string]string) (*types.Rewrite, error) {
rewritePayload := annotations[types.DiscoveryAppRewriteLabel]
if rewritePayload == "" {
return nil, nil
}
rw := types.Rewrite{}
reader := strings.NewReader(rewritePayload)
decoder := kyaml.NewYAMLOrJSONDecoder(reader, 32*1024)
err := decoder.Decode(&rw)
if err != nil {
return nil, trace.Wrap(err, "failed decoding rewrite config")
}
return &rw, nil
}
func getDescription(annotations map[string]string) string {
return annotations[types.DiscoveryDescription]
}
func getPublicAddr(annotations map[string]string) string {
return annotations[types.DiscoveryPublicAddr]
}
func getTLSInsecureSkipVerify(annotations map[string]string) bool {
val := annotations[types.DiscoveryAppInsecureSkipVerify]
if val == "" {
return false
}
return val == "true"
}
func getAppName(serviceName, namespace, clusterName, portName, nameAnnotation string) (string, error) {
if nameAnnotation != "" {
name := nameAnnotation
if portName != "" {
name = fmt.Sprintf("%s-%s", name, portName)
}
if len(validation.IsDNS1123Label(name)) > 0 {
return "", trace.BadParameter(
"application name %q must be a valid DNS label (lowercase alphanumeric or '-', must start and end with alphanumeric, max 63 chars): https://goteleport.com/docs/enroll-resources/application-access/guides/connecting-apps/#application-name", name)
}
return name, nil
}
// Lowercase + dot-replace so the composed name passes ValidateApp.
clusterName = strings.ToLower(strings.ReplaceAll(clusterName, ".", "-"))
if portName != "" {
return fmt.Sprintf("%s-%s-%s-%s", serviceName, portName, namespace, clusterName), nil
}
return fmt.Sprintf("%s-%s-%s", serviceName, namespace, clusterName), nil
}
func getAppLabels(serviceLabels map[string]string, clusterName string) (map[string]string, error) {
result := make(map[string]string, len(serviceLabels)+1)
for k, v := range serviceLabels {
if !types.IsValidLabelKey(k) {
return nil, trace.BadParameter("invalid label key: %q", k)
}
result[k] = v
}
result[types.KubernetesClusterLabel] = clusterName
return result, nil
}
var (
// clusterDomainResolver is a function that resolves the cluster domain once and caches the result.
// It's used to lazily resolve the cluster domain from the env var "TELEPORT_KUBE_CLUSTER_DOMAIN" or fallback to
// a default value.
// It's only used when agent is running in the Kubernetes cluster.
clusterDomainResolver = sync.OnceValue[string](getClusterDomain)
)
const (
// teleportKubeClusterDomain is the environment variable that specifies the cluster domain.
teleportKubeClusterDomain = "TELEPORT_KUBE_CLUSTER_DOMAIN"
)
func getClusterDomain() string {
if envDomain := os.Getenv(teleportKubeClusterDomain); envDomain != "" {
return envDomain
}
return "cluster.local"
}
// RewriteHeadersAndApplyValueTraits rewrites the provided request's headers
// while applying value traits to them.
func RewriteHeadersAndApplyValueTraits(r *http.Request, rewrites iter.Seq[*types.Header], rewriteTraits wrappers.Traits, log *slog.Logger) {
for header := range rewrites {
values, err := ApplyValueTraits(header.Value, rewriteTraits)
if err != nil {
log.DebugContext(r.Context(), "Failed to apply traits",
"header_value", header.Value,
"error", err,
)
continue
}
r.Header.Del(header.Name)
for _, value := range values {
switch http.CanonicalHeaderKey(header.Name) {
case teleport.HostHeader:
r.Host = value
default:
r.Header.Add(header.Name, value)
}
}
}
}
// isValidSpiffeID validates that s contains a valid SPIFFE ID.
func isValidSpiffeID(s string) error {
_, err := spiffeid.FromString(s)
return err
}
// isValidCACertificatePEM validates that s contains valid PEM-encoded CA
// certificate.
func isValidCACertificatePEM(s string) error {
cert, err := tlsutils.ParseCertificatePEMStrict([]byte(s))
if err != nil {
return trace.Wrap(err)
}
switch {
case !cert.BasicConstraintsValid || !cert.IsCA:
return trace.BadParameter("certificate %q is not a CA", cert.Subject.String())
case cert.KeyUsage != 0 && cert.KeyUsage&x509.KeyUsageCertSign == 0:
return trace.BadParameter("CA certificate %q does not allow certificate signing", cert.Subject.String())
}
return nil
}
// quoteAndJoin takes a slice of strings and returns them quoted and
// comma-separated.
func quoteAndJoin(items []string) string {
if len(items) == 0 {
return ""
}
quotedItems := make([]string, len(items))
for i, item := range items {
quotedItems[i] = `"` + item + `"`
}
return strings.Join(quotedItems, ", ")
}
/*
* Teleport
* Copyright (C) 2026 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 (
"time"
"google.golang.org/protobuf/proto"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types"
)
// AppAccessChecker provides app-specific access checking, abstracting over scoped and unscoped
// identities. It is obtained from [ScopedAccessChecker.App] and should not be constructed directly.
// Methods on this type implement app-specific behavior, branching internally between the scoped and
// unscoped paths of the underlying [ScopedAccessChecker].
type AppAccessChecker struct {
checker *ScopedAccessChecker
}
// CheckAccessToApp checks access to an application.
func (c *AppAccessChecker) CheckAccessToApp(target types.Application, state AccessState, matchers ...RoleMatcher) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAccess(target, state, matchers...)
}
return c.checker.scopedCompatChecker.CheckAccess(target, state, matchers...)
}
// CanAccessApp checks whether read access to the specified application is possible without regard
// to a specific MFA state. Used for listing/filtering.
func (c *AppAccessChecker) CanAccessApp(target types.Application) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAccess(target, AccessState{MFAVerified: true})
}
return c.checker.scopedCompatChecker.CheckAccess(target, AccessState{MFAVerified: true})
}
// AdjustClientIdleTimeout determines the app client idle timeout to apply.
func (c *AppAccessChecker) AdjustClientIdleTimeout(timeout time.Duration) (time.Duration, error) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.AdjustClientIdleTimeout(timeout), nil
}
return c.checker.adjustScopedClientIdleTimeout(c.checker.role.GetSpec().GetApp().GetClientIdleTimeout(), timeout)
}
// AdjustDisconnectExpiredCert adjusts whether to disconnect on certificate expiry.
func (c *AppAccessChecker) AdjustDisconnectExpiredCert(disconnect bool) bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.AdjustDisconnectExpiredCert(disconnect)
}
app := c.checker.role.GetSpec().GetApp()
var disconnectExpiredCert *bool
if app != nil {
disconnectExpiredCert = proto.ValueOrNil(app.HasDisconnectExpiredCert(), app.GetDisconnectExpiredCert)
}
return c.checker.adjustScopedDisconnectExpiredCert(disconnectExpiredCert, disconnect)
}
// LockingMode returns the App lock enforcement mode to apply.
func (c *AppAccessChecker) LockingMode(defaultMode constants.LockingMode) constants.LockingMode {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.LockingMode(defaultMode)
}
return c.checker.scopedLockingMode(c.checker.role.GetSpec().GetApp().GetLock(), defaultMode)
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/utils"
)
// ClusterAuditConfigSpecFromObject returns audit config spec from object.
func ClusterAuditConfigSpecFromObject(in any) (*types.ClusterAuditConfigSpecV2, error) {
var cfg types.ClusterAuditConfigSpecV2
if in == nil {
return &cfg, nil
}
if err := apiutils.ObjectToStruct(in, &cfg); err != nil {
return nil, trace.Wrap(err)
}
return &cfg, nil
}
// UnmarshalClusterAuditConfig unmarshals the ClusterAuditConfig resource from JSON.
func UnmarshalClusterAuditConfig(bytes []byte, opts ...MarshalOption) (types.ClusterAuditConfig, error) {
var auditConfig types.ClusterAuditConfigV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &auditConfig); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := auditConfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
auditConfig.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
auditConfig.SetExpiry(cfg.Expires)
}
return &auditConfig, nil
}
// MarshalClusterAuditConfig marshals the ClusterAuditConfig resource to JSON.
func MarshalClusterAuditConfig(auditConfig types.ClusterAuditConfig, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch auditConfig := auditConfig.(type) {
case *types.ClusterAuditConfigV2:
if err := auditConfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, auditConfig))
default:
return nil, trace.BadParameter("unrecognized cluster audit config version %T", auditConfig)
}
}
/*
* 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 types contains all types and logic required by the Teleport API.
package services
import (
"time"
"github.com/google/uuid"
"github.com/gravitational/trace"
"golang.org/x/crypto/bcrypt"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// ValidateLocalAuthSecrets validates local auth secret members.
func ValidateLocalAuthSecrets(l *types.LocalAuthSecrets) error {
if len(l.PasswordHash) > 0 {
if _, err := bcrypt.Cost(l.PasswordHash); err != nil {
return trace.BadParameter("invalid password hash")
}
}
mfaNames := make(map[string]struct{}, len(l.MFA))
for _, d := range l.MFA {
if err := d.CheckAndSetDefaults(); err != nil {
return trace.BadParameter("MFA device named %q is invalid: %v", d.Metadata.Name, err)
}
if _, ok := mfaNames[d.Metadata.Name]; ok {
return trace.BadParameter("MFA device named %q already exists", d.Metadata.Name)
}
mfaNames[d.Metadata.Name] = struct{}{}
}
if l.Webauthn != nil {
if err := l.Webauthn.Check(); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// NewTOTPDevice creates a TOTP MFADevice from the given key.
func NewTOTPDevice(name, key string, addedAt time.Time) (*types.MFADevice, error) {
d, err := types.NewMFADevice(name, uuid.New().String(), addedAt, &types.MFADevice_Totp{Totp: &types.TOTPDevice{
Key: key,
}})
if err != nil {
return nil, trace.Wrap(err)
}
return d, nil
}
// UnmarshalAuthPreference unmarshals the AuthPreference resource from JSON.
func UnmarshalAuthPreference(bytes []byte, opts ...MarshalOption) (types.AuthPreference, error) {
var authPreference types.AuthPreferenceV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &authPreference); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := authPreference.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
authPreference.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
authPreference.SetExpiry(cfg.Expires)
}
return &authPreference, nil
}
// MarshalAuthPreference marshals the AuthPreference resource to JSON.
func MarshalAuthPreference(c types.AuthPreference, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch c := c.(type) {
case *types.AuthPreferenceV2:
if err := c.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *c
copy.SetRevision("")
c = ©
}
return utils.FastMarshal(c)
default:
return nil, trace.BadParameter("unsupported type for auth preference: %T", c)
}
}
/*
* 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 (
"crypto"
"crypto/tls"
"crypto/x509"
"iter"
"slices"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/lib/jwt"
"github.com/gravitational/teleport/lib/sshutils"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
)
// ValidateCertAuthority validates the CertAuthority
func ValidateCertAuthority(ca types.CertAuthority) (err error) {
if err = CheckAndSetDefaults(ca); err != nil {
return trace.Wrap(err)
}
switch ca.GetType() {
case types.UserCA, types.HostCA:
err = checkUserOrHostCA(ca)
case types.DatabaseCA, types.DatabaseClientCA:
err = checkDatabaseCA(ca)
case types.OpenSSHCA:
err = checkOpenSSHCA(ca)
case types.JWTSigner, types.OIDCIdPCA, types.OktaCA, types.BoundKeypairCA:
err = checkJWTKeys(ca)
case types.SAMLIDPCA:
err = checkSAMLIDPCA(ca)
case types.SPIFFECA:
err = checkSPIFFECA(ca)
case types.AWSRACA:
err = checkAWSRACA(ca)
case types.WindowsCA:
err = checkWindowsCA(ca)
case types.AppClientCA:
err = checkAppClientCA(ca)
default:
return trace.BadParameter("invalid CA type %q", ca.GetType())
}
return trace.Wrap(err)
}
func checkSPIFFECA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
if len(ca.Spec.ActiveKeys.TLS) == 0 {
return trace.BadParameter("certificate authority missing TLS key pairs")
}
if len(ca.Spec.ActiveKeys.JWT) == 0 {
return trace.BadParameter("certificate authority missing JWT key pairs")
}
return nil
}
func checkAWSRACA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
if len(ca.Spec.ActiveKeys.TLS) == 0 {
return trace.BadParameter("certificate authority missing TLS key pairs")
}
return nil
}
func checkUserOrHostCA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
if len(ca.Spec.ActiveKeys.SSH) == 0 {
return trace.BadParameter("certificate authority missing SSH key pairs")
}
if len(ca.Spec.ActiveKeys.TLS) == 0 {
return trace.BadParameter("certificate authority missing TLS key pairs")
}
if _, err := sshutils.GetCheckers(ca); err != nil {
return trace.Wrap(err)
}
if err := sshutils.ValidateSigners(ca); err != nil {
return trace.Wrap(err)
}
// This is to force users to migrate
if len(ca.GetRoles()) != 0 && len(ca.GetRoleMap()) != 0 {
return trace.BadParameter("should set either 'roles' or 'role_map', not both")
}
_, err := parseRoleMap(ca.GetRoleMap())
return trace.Wrap(err)
}
// checkDatabaseCA checks if provided certificate authority contains a valid TLS key pair.
// This function is used to verify Database CA.
func checkDatabaseCA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
return trace.Wrap(checkTLSKeys(ca))
}
// checkOpenSSHCA checks if provided certificate authority contains a valid SSH key pair.
func checkOpenSSHCA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
if len(ca.Spec.ActiveKeys.SSH) == 0 {
return trace.BadParameter("certificate authority missing SSH key pairs")
}
if _, err := sshutils.GetCheckers(ca); err != nil {
return trace.Wrap(err)
}
if err := sshutils.ValidateSigners(ca); err != nil {
return trace.Wrap(err)
}
// This is to force users to migrate
if len(ca.GetRoles()) != 0 && len(ca.GetRoleMap()) != 0 {
return trace.BadParameter("should set either 'roles' or 'role_map', not both")
}
_, err := parseRoleMap(ca.GetRoleMap())
return trace.Wrap(err)
}
func checkJWTKeys(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
// Check that some JWT keys have been set on the CA.
if len(ca.Spec.ActiveKeys.JWT) == 0 {
return trace.BadParameter("missing JWT CA")
}
var err error
var privateKey crypto.Signer
// Check that the JWT keys set are valid.
for _, pair := range ca.GetTrustedJWTKeyPairs() {
// TODO(nic): validate PKCS11 private keys
if len(pair.PrivateKey) > 0 && pair.PrivateKeyType == types.PrivateKeyType_RAW {
privateKey, err = keys.ParsePrivateKey(pair.PrivateKey)
if err != nil {
return trace.Wrap(err)
}
}
publicKey, err := keys.ParsePublicKey(pair.PublicKey)
if err != nil {
return trace.Wrap(err)
}
cfg := &jwt.Config{
ClusterName: ca.GetClusterName(),
PrivateKey: privateKey,
PublicKey: publicKey,
}
if _, err = jwt.New(cfg); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// checkSAMLIDPCA checks if provided certificate authority contains a valid TLS key pair.
// This function is used to verify the SAML IDP CA.
func checkSAMLIDPCA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
return trace.Wrap(checkTLSKeys(ca))
}
func checkWindowsCA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
return trace.Wrap(checkTLSKeys(ca))
}
func checkAppClientCA(cai types.CertAuthority) error {
ca, ok := cai.(*types.CertAuthorityV2)
if !ok {
return trace.BadParameter("unknown CA type %T", cai)
}
return trace.Wrap(checkTLSKeys(ca))
}
func checkTLSKeys(ca *types.CertAuthorityV2) error {
if len(ca.Spec.ActiveKeys.TLS) == 0 {
return trace.BadParameter("%s certificate authority missing TLS key pairs", ca.GetType())
}
for _, pair := range ca.GetTrustedTLSKeyPairs() {
// Note: A non-empty pair.Cert is required by pair.CheckAndSetDefaults().
if len(pair.Key) > 0 && pair.KeyType == types.PrivateKeyType_RAW {
if _, err := tls.X509KeyPair(pair.Cert, pair.Key); err != nil {
return trace.Wrap(err, "private key and certificate")
}
continue
}
if _, err := tlsca.ParseCertificatePEM(pair.Cert); err != nil {
return trace.Wrap(err, "certificate")
}
}
return nil
}
// GetJWTSigner returns the active JWT key used to sign tokens.
func GetJWTSigner(signer crypto.Signer, clusterName string, clock clockwork.Clock) (*jwt.Key, error) {
key, err := jwt.New(&jwt.Config{
Clock: clock,
ClusterName: clusterName,
PrivateKey: signer,
})
return key, trace.Wrap(err)
}
// GetTLSCerts returns TLS certificates from CA
func GetTLSCerts(ca types.CertAuthority) [][]byte {
pairs := ca.GetTrustedTLSKeyPairs()
out := make([][]byte, len(pairs))
for i, pair := range pairs {
out[i] = slices.Clone(pair.Cert)
}
return out
}
// GetX509Certs returns parsed TLS certificates from CA as [x509.Certificate].
func GetX509Certs(ca types.CertAuthority) iter.Seq2[*x509.Certificate, error] {
pairs := ca.GetTrustedTLSKeyPairs()
return func(yield func(*x509.Certificate, error) bool) {
for _, pair := range pairs {
cert, err := tlsca.ParseCertificatePEM(pair.Cert)
if !yield(cert, err) {
return
}
}
}
}
// GetSSHCheckingKeys returns SSH public keys from CA
func GetSSHCheckingKeys(ca types.CertAuthority) [][]byte {
pairs := ca.GetTrustedSSHKeyPairs()
out := make([][]byte, 0, len(pairs))
for _, pair := range pairs {
out = append(out, slices.Clone(pair.PublicKey))
}
return out
}
// CertPoolFromCertAuthorities returns a certificate pool from the TLS certificates
// set up in the certificate authorities list, as well as the number of certificates
// that were added to the pool.
func CertPoolFromCertAuthorities(cas []types.CertAuthority) (*x509.CertPool, int, error) {
certPool := x509.NewCertPool()
count := 0
for _, ca := range cas {
for cert, err := range GetX509Certs(ca) {
if err != nil {
return nil, 0, trace.Wrap(err)
}
certPool.AddCert(cert)
count++
}
}
return certPool, count, nil
}
// CertPool returns certificate pools from TLS certificates
// set up in the certificate authority
func CertPool(ca types.CertAuthority) (*x509.CertPool, error) {
certPool, count, err := CertPoolFromCertAuthorities([]types.CertAuthority{ca})
if err != nil {
return nil, trace.Wrap(err)
}
if count == 0 {
return nil, trace.BadParameter("certificate authority has no TLS certificates")
}
return certPool, nil
}
// UnmarshalCertAuthority unmarshals the CertAuthority resource to JSON.
func UnmarshalCertAuthority(bytes []byte, opts ...MarshalOption) (types.CertAuthority, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
err = utils.FastUnmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V2:
var ca types.CertAuthorityV2
if err := utils.FastUnmarshal(bytes, &ca); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := ValidateCertAuthority(&ca); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
ca.SetRevision(cfg.Revision)
}
// Correct problems with existing CAs that contain non-UTC times, which
// causes panics when doing a gogoproto Clone; should only ever be
// possible with LastRotated, but we enforce it on all the times anyway.
// See https://github.com/gogo/protobuf/issues/519 .
if ca.Spec.Rotation != nil {
apiutils.UTC(&ca.Spec.Rotation.Started)
apiutils.UTC(&ca.Spec.Rotation.LastRotated)
apiutils.UTC(&ca.Spec.Rotation.Schedule.UpdateClients)
apiutils.UTC(&ca.Spec.Rotation.Schedule.UpdateServers)
apiutils.UTC(&ca.Spec.Rotation.Schedule.Standby)
}
return &ca, nil
}
return nil, trace.BadParameter("cert authority resource version %v is not supported", h.Version)
}
// MarshalCertAuthority marshals the CertAuthority resource to JSON.
func MarshalCertAuthority(certAuthority types.CertAuthority, opts ...MarshalOption) ([]byte, error) {
if err := ValidateCertAuthority(certAuthority); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch certAuthority := certAuthority.(type) {
case *types.CertAuthorityV2:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, certAuthority))
default:
return nil, trace.BadParameter("unrecognized certificate authority version %T", certAuthority)
}
}
/*
* Teleport
* Copyright (C) 2026 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"
"iter"
"regexp"
"strings"
"github.com/gravitational/trace"
"k8s.io/apimachinery/pkg/util/validation"
beamsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/beams/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/utils/set"
)
var beamAliasRegexp = regexp.MustCompile(`^[a-z]+-[a-z]+$`)
// BeamReader defines methods for reading beam resources.
type BeamReader interface {
// GetBeam fetches a beam by name.
GetBeam(ctx context.Context, name string) (*beamsv1.Beam, error)
// GetBeamByAlias fetches a beam by alias.
GetBeamByAlias(ctx context.Context, alias string) (*beamsv1.Beam, error)
// ListBeams lists beams with pagination.
ListBeams(ctx context.Context, limit int, startKey string) ([]*beamsv1.Beam, string, error)
// ListBeamsV2 lists beams with pagination, sorting and filtering.
ListBeamsV2(ctx context.Context, limit int, startKey string, options *ListBeamsRequestOptions) ([]*beamsv1.Beam, string, error)
// IterateBeams returns a sequence of beams starting from the given
// pageToken.
IterateBeams(ctx context.Context, pageToken string) iter.Seq2[*beamsv1.Beam, error]
// IterateBeamsV2 returns a sequence of beams starting from the given
// pageToken with sorting and filtering.
IterateBeamsV2(ctx context.Context, pageToken string, options *ListBeamsRequestOptions) iter.Seq2[*beamsv1.Beam, error]
}
// BeamWriter defines methods for writing beam resources. We always write beams
// using Backend.AtomicWrite (with their supporting resources) so this interface
// doesn't contain the usual CRUD methods.
type BeamWriter interface {
// AppendPutBeamActions adds conditional actions to an atomic write to create
// or update a Beam resource.
AppendPutBeamActions(
actions []backend.ConditionalAction,
beam *beamsv1.Beam,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteBeamActions adds conditional actions to an atomic write to
// delete a Beam resource.
AppendDeleteBeamActions(
actions []backend.ConditionalAction,
beam *beamsv1.Beam,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
// Beams defines methods for managing beam resources.
type Beams interface {
BeamReader
BeamWriter
}
// ValidateBeam validates the given beam resource.
func ValidateBeam(b *beamsv1.Beam) error {
switch {
case b == nil:
return trace.BadParameter("beam must not be nil")
case b.GetVersion() != types.V1:
return trace.BadParameter("version: only supports version %q, got %q", types.V1, b.GetVersion())
case b.GetKind() != types.KindBeam:
return trace.BadParameter("kind: must be %q, got %q", types.KindBeam, b.GetKind())
case !b.HasMetadata():
return trace.BadParameter("metadata: is required")
case b.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name: is required")
case !b.HasSpec():
return trace.BadParameter("spec: is required")
case !b.GetSpec().HasExpires():
return trace.BadParameter("spec.expires: is required")
case !b.HasStatus():
return trace.BadParameter("status: is required")
}
switch b.GetSpec().GetEgress() {
case beamsv1.EgressMode_EGRESS_MODE_RESTRICTED:
for i, domain := range b.GetSpec().GetAllowedDomains() {
// Must be fully-qualified, up to the root.
if !strings.HasSuffix(domain, ".") {
return trace.BadParameter("spec.allowed_domains[%d]: %q must be a fully qualified domain name ending with '.'", i, domain)
}
trimmedDomain := strings.TrimSuffix(domain, ".")
// TLDs like "com." or "net." are invalid, as is "localhost."
if !strings.Contains(trimmedDomain, ".") {
return trace.BadParameter("spec.allowed_domains[%d]: %q must be a fully qualified domain name ending with '.'", i, domain)
}
// Note: wildcard like "*.example.com." are explicitly not supported
// because we don't yet have a way to proxy them via VNet.
if errs := validation.IsDNS1123Subdomain(trimmedDomain); len(errs) > 0 {
return trace.BadParameter("spec.allowed_domains[%d]: %q is invalid: %s", i, domain, errs)
}
}
case beamsv1.EgressMode_EGRESS_MODE_UNRESTRICTED:
if len(b.GetSpec().GetAllowedDomains()) > 0 {
return trace.BadParameter("spec.allowed_domains: may only be set when spec.egress is EGRESS_MODE_RESTRICTED")
}
default:
return trace.BadParameter("spec.egress: must be EGRESS_MODE_RESTRICTED or EGRESS_MODE_UNRESTRICTED, got %s", b.GetSpec().GetEgress())
}
if pub := b.GetSpec().GetPublish(); pub != nil {
if pub.GetPort() != 8080 {
return trace.BadParameter("spec.publish.port: must be 8080")
}
switch pub.GetProtocol() {
case beamsv1.Protocol_PROTOCOL_HTTP, beamsv1.Protocol_PROTOCOL_TCP:
default:
return trace.BadParameter("spec.publish.protocol: must be HTTP or TCP")
}
}
if err := ValidateBeamAlias(b.GetStatus().GetAlias()); err != nil {
return trace.Wrap(err)
}
return nil
}
func ValidateBeamAlias(alias string) error {
if beamAliasRegexp.MatchString(alias) {
return nil
}
return trace.BadParameter("beam alias must be a hyphen-separated pair of two lowercase words")
}
// MakeBeamFilterFunc creates a filter function for beams based on the provided
// options.
func MakeBeamFilterFunc(options *ListBeamsRequestOptions) func(beam *beamsv1.Beam) bool {
return func(b *beamsv1.Beam) bool {
if options.GetFilterUsers().Len() > 0 && !options.FilterUsers.Contains(b.GetStatus().GetUser()) {
return false
}
if options.GetFilterFn() != nil && !options.FilterFn(b) {
return false
}
return true
}
}
type ListBeamsRequestOptions struct {
// The sort field to use for the results. If unspecified, the default sort
// field is used.
SortField beamsv1.BeamSortField
// The sort order to use for the results. If unspecified, the default sort
// order is used.
SortOrder beamsv1.BeamSortOrder
// FilterUsers filters the results to only include beams owned by the
// provided users.
FilterUsers set.Set[string]
// FilterFn is a general-use filter delegate. Useful when the state required
// for a filter means the filter can't be easily implemented in the backend
// or cache (e.g. access control context).
FilterFn func(*beamsv1.Beam) bool
}
// GetSortField is a nil-safe getter for SortField
func (o *ListBeamsRequestOptions) GetSortField() beamsv1.BeamSortField {
if o == nil {
return beamsv1.BeamSortField_BEAM_SORT_FIELD_UNSPECIFIED
}
return o.SortField
}
// GetSortOrder is a nil-safe getter for SortDesc
func (o *ListBeamsRequestOptions) GetSortOrder() beamsv1.BeamSortOrder {
if o == nil {
return beamsv1.BeamSortOrder_BEAM_SORT_ORDER_UNSPECIFIED
}
return o.SortOrder
}
// GetFilterUsers is a nil-safe getter for FilterOwners
func (o *ListBeamsRequestOptions) GetFilterUsers() set.Set[string] {
if o == nil {
return set.Set[string]{}
}
return o.FilterUsers
}
// GetFilterFn is a nil-safe getter for FilterFn
func (o *ListBeamsRequestOptions) GetFilterFn() func(*beamsv1.Beam) bool {
if o == nil {
return nil
}
return o.FilterFn
}
// Teleport
// Copyright (C) 2026 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"
beamsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/beams/v1"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
"github.com/gravitational/teleport/api/types"
)
// BeamsConfigGetter is an interface for getting the user-provided BeamsConfig singleton.
type BeamsConfigGetter interface {
// GetBeamsConfig returns the singleton BeamsConfig resource.
GetBeamsConfig(ctx context.Context) (*beamsv1.BeamsConfig, error)
}
// BeamsConfigService is an interface for managing the user-provided BeamsConfig singleton.
type BeamsConfigService interface {
BeamsConfigGetter
// CreateBeamsConfig creates a new BeamsConfig resource.
CreateBeamsConfig(ctx context.Context, config *beamsv1.BeamsConfig) (*beamsv1.BeamsConfig, error)
// UpdateBeamsConfig updates an existing BeamsConfig resource using conditional update.
UpdateBeamsConfig(ctx context.Context, config *beamsv1.BeamsConfig) (*beamsv1.BeamsConfig, error)
// DeleteBeamsConfig deletes the singleton BeamsConfig resource.
DeleteBeamsConfig(ctx context.Context) error
}
// DefaultBeamsConfig returns the default virtual BeamsConfig resource.
// App names default to the cloud-managed LLM app names. These are not stored
// in the backend and are returned when no user-created resource exists.
func DefaultBeamsConfig() *beamsv1.BeamsConfig {
return beamsv1.BeamsConfig_builder{
Kind: types.KindBeamsConfig,
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: types.MetaNameBeamsConfig,
}.Build(),
Spec: beamsv1.BeamsConfigSpec_builder{
Llm: beamsv1.LLMConfig_builder{
Anthropic: beamsv1.LLMEndpointConfig_builder{
AppName: "anthropic",
}.Build(),
Openai: beamsv1.LLMEndpointConfig_builder{
AppName: "openai",
}.Build(),
}.Build(),
}.Build(),
}.Build()
}
// Teleport
// Copyright (C) 2024 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"
"slices"
"strings"
"github.com/charlievieth/strcase"
"github.com/gravitational/trace"
machineidv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/lib/auth/machineid/machineidv1/expression"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils/typical"
)
// BotUserPrefix is the prefix appended to bot users. Users with this prefix are
// plausibly bots, but this prefix is not sufficient to establish that a given
// user actually is a bot user. The existence of [types.BotLabel] in the user
// labels is the definitive indicator that a user is a bot.
const BotUserPrefix = "bot-"
// BotInstance is an interface for the BotInstance service.
//
// Bot instances belong to a bot, and bots are identified by their scope and
// name: instances of scoped bots are stored in a scope-namespaced key range,
// separate from instances of unscoped bots. Methods addressing an individual
// instance take the owning bot's scope, which must be empty if the bot is
// unscoped.
type BotInstance interface {
// CreateBotInstance creates a new bot instance. It is stored in the key
// range determined by the scope set on the instance itself.
CreateBotInstance(ctx context.Context, botInstance *machineidv1.BotInstance) (*machineidv1.BotInstance, error)
// GetBotInstance returns the bot instance owned by the bot identified by
// the request's (bot_scope, bot_name) with the given instance ID.
GetBotInstance(ctx context.Context, req *machineidv1.GetBotInstanceRequest) (*machineidv1.BotInstance, error)
// ListBotInstances
ListBotInstances(ctx context.Context, pageSize int, lastToken string, options *ListBotInstancesRequestOptions) ([]*machineidv1.BotInstance, string, error)
// DeleteBotInstance deletes the bot instance owned by the bot identified
// by the request's (bot_scope, bot_name) with the given instance ID.
DeleteBotInstance(ctx context.Context, req *machineidv1.DeleteBotInstanceRequest) error
// PatchBotInstance fetches the existing bot instance identified by the
// given options, then calls the options' UpdateFn to apply any changes
// before persisting the resource.
PatchBotInstance(ctx context.Context, opts PatchBotInstanceOpts) (*machineidv1.BotInstance, error)
}
// PatchBotInstanceOpts identifies the bot instance to be patched by
// [BotInstance.PatchBotInstance] and holds the patch to apply to it.
type PatchBotInstanceOpts struct {
// Bot is the scope-qualified name of the bot that owns the instance. The
// scope must be empty if the bot is unscoped.
Bot scopes.QualifiedName
// InstanceID is the ID of the instance to patch.
InstanceID string
// UpdateFn is applied to the fetched instance to produce the instance to
// persist. It may be called more than once if the write is retried.
UpdateFn func(*machineidv1.BotInstance) (*machineidv1.BotInstance, error)
}
// ValidateBotInstance verifies that required fields for a new BotInstance are present
func ValidateBotInstance(b *machineidv1.BotInstance) error {
if !b.HasSpec() {
return trace.BadParameter("spec is required")
}
if b.GetSpec().GetBotName() == "" {
return trace.BadParameter("spec.bot_name is required")
}
if b.GetSpec().GetInstanceId() == "" {
return trace.BadParameter("spec.instance_id is required")
}
if !b.HasStatus() {
return trace.BadParameter("status is required")
}
return nil
}
// MarshalBotInstance marshals the BotInstance object into a JSON byte array.
func MarshalBotInstance(object *machineidv1.BotInstance, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalBotInstance unmarshals the BotInstance object from a JSON byte array.
func UnmarshalBotInstance(data []byte, opts ...MarshalOption) (*machineidv1.BotInstance, error) {
return UnmarshalProtoResource[*machineidv1.BotInstance](data, opts...)
}
func MatchBotInstance(b *machineidv1.BotInstance, botName string, search string, exp typical.Expression[*expression.Environment, bool]) bool {
if botName != "" && b.GetSpec().GetBotName() != botName {
return false
}
heartbeat := GetBotInstanceLatestHeartbeat(b)
authentication := GetBotInstanceLatestAuthentication(b)
if exp != nil {
if match, err := exp.Evaluate(&expression.Environment{
Metadata: b.GetMetadata(),
Spec: b.GetSpec(),
LatestHeartbeat: heartbeat,
LatestAuthentication: authentication,
}); err != nil || !match {
return false
}
}
if search == "" {
return true
}
values := []string{
b.GetSpec().GetBotName(),
b.GetSpec().GetInstanceId(),
}
if heartbeat != nil {
values = append(values, heartbeat.GetHostname(), heartbeat.GetJoinMethod(), heartbeat.GetVersion(), "v"+heartbeat.GetVersion())
}
return slices.ContainsFunc(values, func(val string) bool {
return strcase.Contains(val, search)
})
}
// GetBotInstanceLatestHeartbeat returns the most recent heartbeat for the
// given bot instance.
func GetBotInstanceLatestHeartbeat(botInstance *machineidv1.BotInstance) *machineidv1.BotInstanceStatusHeartbeat {
heartbeat := botInstance.GetStatus().GetInitialHeartbeat()
latestHeartbeats := botInstance.GetStatus().GetLatestHeartbeats()
if len(latestHeartbeats) > 0 {
heartbeat = latestHeartbeats[len(latestHeartbeats)-1]
}
return heartbeat
}
// GetBotInstanceLatestAuthentication returns the most recent authentication for
// the given bot instance.
func GetBotInstanceLatestAuthentication(botInstance *machineidv1.BotInstance) *machineidv1.BotInstanceStatusAuthentication {
authentication := botInstance.GetStatus().GetInitialAuthentication()
latestAuthentications := botInstance.GetStatus().GetLatestAuthentications()
if len(latestAuthentications) > 0 {
authentication = latestAuthentications[len(latestAuthentications)-1]
}
return authentication
}
type ListBotInstancesRequestOptions struct {
// The sort field to use for the results. If empty, the default sort field
// is used.
SortField string
// The sort order to use for the results. If empty, the default sort order
// is used.
SortDesc bool
// The name of the Bot to list BotInstances for. If empty, all BotInstances
// will be listed.
FilterBotName string
// The scope of the Bot to list BotInstances for. A bot is identified by the
// pair (scope, name), so this only ever qualifies FilterBotName and must be
// set alongside it; setting it without FilterBotName is an error. Leave
// empty if the bot is unscoped. This is deliberately not a scope filter for
// listing every BotInstance in a scope - use ScopeFilter for that.
FilterBotScope string
// ScopeFilter selects BotInstances by the scope of their owning bot. A nil or
// MODE_UNSPECIFIED filter matches every scope; identity-derived defaulting is
// the caller's job (see ScopedAccessCheckerContext.ResolveScopeFilter).
//
// Mutually exclusive with FilterBotName, which already constrains the result
// to a single bot in a single scope.
ScopeFilter *scopesv1.Filter
// A search term used to filter the results. If non-empty, it's used to
// match against supported fields.
FilterSearchTerm string
// A Teleport predicate language query used to filter the results.
FilterQuery string
// FilterFn is an optional additional filter applied during iteration.
FilterFn func(*machineidv1.BotInstance) bool
}
func (o *ListBotInstancesRequestOptions) GetSortField() string {
if o == nil {
return ""
}
return o.SortField
}
func (o *ListBotInstancesRequestOptions) GetSortDesc() bool {
if o == nil {
return false
}
return o.SortDesc
}
func (o *ListBotInstancesRequestOptions) GetFilterBotName() string {
if o == nil {
return ""
}
return o.FilterBotName
}
func (o *ListBotInstancesRequestOptions) GetFilterBotScope() string {
if o == nil {
return ""
}
return o.FilterBotScope
}
func (o *ListBotInstancesRequestOptions) GetScopeFilter() *scopesv1.Filter {
if o == nil {
return nil
}
return o.ScopeFilter
}
func (o *ListBotInstancesRequestOptions) GetFilterSearchTerm() string {
if o == nil {
return ""
}
return o.FilterSearchTerm
}
func (o *ListBotInstancesRequestOptions) GetFilterQuery() string {
if o == nil {
return ""
}
return o.FilterQuery
}
func (o *ListBotInstancesRequestOptions) GetFilterFn() func(*machineidv1.BotInstance) bool {
if o == nil {
return nil
}
return o.FilterFn
}
// BotResourceName returns the default name for resources associated with the
// given bot. An empty Scope refers to an unscoped bot.
//
// Bots are namespaced by their scope, so a scoped bot's name encodes the scope
// as well as the bot name (bot-<encoded_scope>-<name>), allowing a name to be reused across
// scopes. An encoded scope only ever contains lowercase alphanumerics, so the "-" separator
// keeps the two apart and two different scopes cannot yield the same name. Scoped bots are
// reconstructed from User labels rather than by parsing this name, so it serves
// only as an identity key.
func BotResourceName(bot scopes.QualifiedName) (string, error) {
name := bot.Name
if bot.Scope != "" {
encodedScope, err := scopes.EncodeForKey(bot.Scope)
if err != nil {
return "", trace.Wrap(err, "encoding scope for bot resource name")
}
name = encodedScope + "-" + bot.Name
}
return BotUserPrefix + strings.ReplaceAll(name, " ", "-"), nil
}
/*
* Teleport
* Copyright (C) 2026 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"
"slices"
"strings"
"time"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/keys"
scopedaccess "github.com/gravitational/teleport/lib/scopes/access"
)
// UnscopedCertificateParameters represents a subset of the AccessChecker interface that
// is used during certificate generation to obtain certificate parameters that are only
// meaningful for unscoped identities.
type UnscopedCertificateParameters interface {
RoleNames() []string
CertificateFormat() string
CertificateExtensions() []*types.CertExtension
CheckKubeGroupsAndUsers(ttl time.Duration, overrideTTL bool, matchers ...RoleMatcher) ([]string, []string, error)
CheckDatabaseNamesAndUsers(ttl time.Duration, overrideTTL bool) ([]string, []string, error)
CheckAWSRoleARNs(ttl time.Duration, overrideTTL bool) ([]string, error)
CheckAzureIdentities(ttl time.Duration, overrideTTL bool) ([]string, error)
CheckGCPServiceAccounts(ttl time.Duration, overrideTTL bool) ([]string, error)
GetAllowedResourceAccessIDs() []types.ResourceAccessID
CheckAccessToRemoteCluster(rc types.RemoteCluster) error
}
// CertificateParameterContext provides methods for resolving certificate parameters that abstract
// over scoped and unscoped identities. Methods on this type should only be called during certificate
// generation and return parameters that need to be embedded in the certificate at issuance time. For
// unscoped identities these parameters are generally equivalent to those returned by the underlying
// AccessChecker. For scoped identities things get more complex as most certificate parameters cannot
// be determined by scoped roles. Instead, parameters for scoped identities are generally hard-coded for
// the time being, with the intent to revisit them in the future and to provide non-role means of
// configuring them. See the Scopes RFD for more details on how scoped permissions intersect with
// certificate parameters.
type CertificateParameterContext struct {
ctx *ScopedAccessCheckerContext
}
// UnscopedCertParams returns unscoped-specific certificate parameters if this is an unscoped
// identity, or nil if this is a scoped identity. Use this for certificate parameters
// that are only meaningful for unscoped identities (e.g., kube groups, db users).
func (c *CertificateParameterContext) UnscopedCertParams() UnscopedCertificateParameters {
return c.ctx.unscopedChecker
}
// GetSSHLoginsForTTL verifies that the requested session TTL is valid and returns
// the list of allowed logins for the certificate.
// - Unscoped: Returns logins from roles, restricted by role TTL rules
// - Scoped: Returns all possible logins across all roles in the pin. this behavior is necessary
// because we cannot determine the effective role without knowing the target resource, but the ssh
// protocol requires all valid principals to be present in the certificate at issuance time. Subsequent
// access checks will enforce login restrictions based on the effective role once the target resource
// is known. Note that this function is *not* safe to determine the logins to be used for OpenSSH agent
// access certs.
func (c *CertificateParameterContext) GetSSHLoginsForTTL(ctx context.Context, ttl time.Duration) ([]string, error) {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.CheckLoginDuration(ttl)
}
// For scoped identities, enumerate all possible logins across all roles in the pin.
// We cannot restrict logins based on a single role since we don't know which role will
// grant access without knowing the target resource.
loginSet := make(map[string]struct{})
// Use of riskyEnumerateScopedCheckers is acceptable here because we are deliberately attempting to aggregate
// information across all roles, rather than making a specific access-control decision.
for checker, err := range c.ctx.riskyEnumerateScopedCheckers(ctx) {
if err != nil {
return nil, trace.Wrap(err)
}
// Get logins from this checker. Pass 0 as TTL to get all logins without TTL restriction.
// We're not enforcing per-role TTL restrictions for scoped certs since the effective role
// is unknown at cert generation time.
for _, login := range checker.SSH().getScopedLogins() {
// Skip placeholder logins when aggregating across roles
if !strings.HasPrefix(login, constants.NoLoginPrefix) {
loginSet[login] = struct{}{}
}
}
}
// Convert map to sorted slice for deterministic output
logins := make([]string, 0, len(loginSet))
for login := range loginSet {
logins = append(logins, login)
}
slices.Sort(logins)
if len(logins) == 0 {
// User was deliberately configured to have no login capability,
// but SSH certificates must contain at least one valid principal.
// We add a single distinctive value which should be unique, and
// will never be a valid unix login (due to leading '-').
logins = []string{constants.NoLoginPrefix + uuid.New().String()}
}
return logins, nil
}
// AdjustSessionTTL adjusts the requested session TTL based on role/configuration policies.
func (c *CertificateParameterContext) AdjustSessionTTL(ttl time.Duration) time.Duration {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.AdjustSessionTTL(ttl)
}
// Scoped identities: return the requested TTL unchanged. We cannot restrict TTL based on roles
// since we don't know which role will grant access without knowing the target resource.
// TODO(fspmarshall/scopes): determine how to handle session TTL restrictions for scoped identities. This will
// likely involve fully decoupling session TTL and certificate TTL, since scoped cert TTLs will need to
// be determined by non-role configuration, whereas specific resource access sessions may still be able to
// be controlled by roles.
return ttl
}
// PrivateKeyPolicy returns the private key policy to enforce for the certificate.
func (c *CertificateParameterContext) PrivateKeyPolicy(defaultPolicy keys.PrivateKeyPolicy) (keys.PrivateKeyPolicy, error) {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.PrivateKeyPolicy(defaultPolicy)
}
// Scoped roles do not currently support custom private key policies. Return the cluster default.
// TODO(fspmarshall/scopes): determine what (if any) control should permit setting the private key
// policy for scoped certificates.
return defaultPolicy, nil
}
// PinSourceIP returns whether source IP pinning should be enabled in the certificate.
func (c *CertificateParameterContext) PinSourceIP() bool {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.PinSourceIP()
}
// Scoped identities do not support source IP pinning due to scope isolation concerns (we can't allow
// to affect certificate parameters).
// TODO(fspmarshall/scopes): determine what (if any) control should permit setting the source IP
// pinning for scoped certificates. Likely this will need to be a cluster configuration rather than
// a role-based setting, though perhapes enablement could be cluster-wide but enforcement could be
// per-role.
return false
}
// CanPortForward returns whether port forwarding should be permitted in the certificate.
func (c *CertificateParameterContext) CanPortForward() bool {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.CanPortForward()
}
// Scoped identities: use unstable env var configuration
// TODO(fspmarshall/scopes): determine what (if any) control should permit setting the port forwarding
// permission for scoped certificates.
return scopedaccess.UnstableGetScopedPortForwarding()
}
// CanForwardAgents returns whether agent forwarding should be permitted in the certificate.
func (c *CertificateParameterContext) CanForwardAgents() bool {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.CanForwardAgents()
}
// Scoped identities: use unstable env var configuration
// TODO(fspmarshall/scopes): determine what (if any) control should permit setting the agent forwarding
// extension for scoped certificates.
return scopedaccess.UnstableGetScopedForwardAgent()
}
// PermitX11Forwarding returns whether X11 forwarding should be permitted in the certificate.
func (c *CertificateParameterContext) PermitX11Forwarding() bool {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.PermitX11Forwarding()
}
// Scoped identities: hard-coded to false (no unstable env var for X11 forwarding)
// TODO(fspmarshall/scopes): determine what (if any) control should permit setting the X11 forwarding
// permission for scoped certificates.
return false
}
// LockingMode returns the locking mode to apply for the certificate.
func (c *CertificateParameterContext) LockingMode(defaultMode constants.LockingMode) constants.LockingMode {
if !c.ctx.isScoped() {
return c.ctx.unscopedChecker.LockingMode(defaultMode)
}
// Scoped roles do not currently support custom locking modes. Return the default/cluster mode.
// TODO(fspmarshall/scopes): determine how to handle locking mode for scoped certificates given that
// role-affected locking behavior during certificate creation doesn't map well to pinned certificates.
return defaultMode
}
/*
* 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 (
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// NewClusterNameWithRandomID creates a ClusterName, supplying a random
// ClusterID if the field is not provided in spec.
func NewClusterNameWithRandomID(spec types.ClusterNameSpecV2) (types.ClusterName, error) {
if spec.ClusterID == "" {
spec.ClusterID = uuid.New().String()
}
return types.NewClusterName(spec)
}
// UnmarshalClusterName unmarshals the ClusterName resource from JSON.
func UnmarshalClusterName(bytes []byte, opts ...MarshalOption) (types.ClusterName, error) {
var clusterName types.ClusterNameV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &clusterName); err != nil {
return nil, trace.BadParameter("%s", err)
}
err = clusterName.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
clusterName.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
clusterName.SetExpiry(cfg.Expires)
}
return &clusterName, nil
}
// MarshalClusterName marshals the ClusterName resource to JSON.
func MarshalClusterName(clusterName types.ClusterName, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch clusterName := clusterName.(type) {
case *types.ClusterNameV2:
if err := clusterName.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, clusterName))
default:
return nil, trace.BadParameter("unrecognized cluster name version %T", clusterName)
}
}
/*
* Teleport
* Copyright (C) 2026 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 (
"time"
"github.com/gravitational/trace"
"google.golang.org/protobuf/proto"
"github.com/gravitational/teleport/api/constants"
accessv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/access/v1"
)
// adjust the client idle timeout value - preferring the default if not set
func (c *ScopedAccessChecker) adjustScopedClientIdleTimeout(idleStr string, timeout time.Duration) (time.Duration, error) {
if idleStr == "" {
idleStr = c.role.GetSpec().GetDefaults().GetClientIdleTimeout()
}
if idleStr != "" {
d, err := time.ParseDuration(idleStr)
if err != nil {
return 0, trace.Errorf("invalid client_idle_timeout %q in scoped role %q: %w", idleStr, c.role.GetMetadata().GetName(), err)
}
if d > 0 && (timeout == 0 || d < timeout) {
return max(d, 0), nil
}
}
return max(timeout, 0), nil
}
// adjustScopedDisconnectExpiredCert returns the disconnect on Expired Cert condition - - applying the default if applicable.
func (c *ScopedAccessChecker) adjustScopedDisconnectExpiredCert(roleSpecifiedDisconnect *bool, defaultdisconnect bool) bool {
if roleSpecifiedDisconnect == nil && c.role.GetSpec().GetDefaults() != nil {
roleSpecifiedDisconnect = proto.ValueOrNil(c.role.GetSpec().GetDefaults().HasDisconnectExpiredCert(), c.role.GetSpec().GetDefaults().GetDisconnectExpiredCert)
}
if roleSpecifiedDisconnect != nil {
return *roleSpecifiedDisconnect
}
return defaultdisconnect
}
// LockingMode returns the lock enforcement mode to apply - applying the default if applicable.
func (c *ScopedAccessChecker) scopedLockingMode(lock *accessv1.Lock, defaultMode constants.LockingMode) constants.LockingMode {
if lock == nil {
lock = c.role.GetSpec().GetDefaults().GetLock()
}
// both protocol specific lock and the default locks are nil, so return the defaultMode.
if lock == nil {
return defaultMode
}
mode := constants.LockingMode(lock.GetMode())
switch mode {
case constants.LockingModeStrict, constants.LockingModeBestEffort:
return mode
default:
return defaultMode
}
}
/*
* 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"
"iter"
"github.com/gravitational/trace"
clusterconfigpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/clusterconfig/v1"
joiningv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/joining/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/modules"
)
// ClusterNameGetter is a service that gets the cluster name from the backend.
type ClusterNameGetter interface {
// GetClusterName gets types.ClusterName from the backend.
GetClusterName(ctx context.Context) (types.ClusterName, error)
}
// ClusterConfiguration stores the cluster configuration in the backend. All
// the resources modified by this interface can only have a single instance
// in the backend.
type ClusterConfiguration interface {
ClusterNameGetter
// GetStaticTokens gets services.StaticTokens from the backend.
GetStaticTokens(context.Context) (types.StaticTokens, error)
// SetStaticTokens sets services.StaticTokens on the backend.
SetStaticTokens(types.StaticTokens) error
// DeleteStaticTokens deletes static tokens resource
DeleteStaticTokens() error
// GetUIConfig gets the proxy service UI config from the backend
GetUIConfig(context.Context) (types.UIConfig, error)
// SetUIConfig sets the proxy service UI config from the backend
SetUIConfig(context.Context, types.UIConfig) error
// DeleteUIConfig deletes the proxy service UI config from the backend
DeleteUIConfig(ctx context.Context) error
// GetAuthPreference gets types.AuthPreference from the backend.
GetAuthPreference(context.Context) (types.AuthPreference, error)
// CreateAuthPreference creates an auth preference if once does not already exist.
CreateAuthPreference(ctx context.Context, preference types.AuthPreference) (types.AuthPreference, error)
// UpdateAuthPreference updates an existing auth preference.
UpdateAuthPreference(ctx context.Context, preference types.AuthPreference) (types.AuthPreference, error)
// UpsertAuthPreference creates a new auth preference or overwrites an existing auth preference.
UpsertAuthPreference(ctx context.Context, preference types.AuthPreference) (types.AuthPreference, error)
// DeleteAuthPreference deletes types.AuthPreference from the backend.
DeleteAuthPreference(ctx context.Context) error
// GetSessionRecordingConfig gets SessionRecordingConfig from the backend.
GetSessionRecordingConfig(context.Context) (types.SessionRecordingConfig, error)
// CreateSessionRecordingConfig creates a session recording config if once does not already exist.
CreateSessionRecordingConfig(ctx context.Context, cfg types.SessionRecordingConfig) (types.SessionRecordingConfig, error)
// UpdateSessionRecordingConfig updates an existing session recording config.
UpdateSessionRecordingConfig(ctx context.Context, cfg types.SessionRecordingConfig) (types.SessionRecordingConfig, error)
// UpsertSessionRecordingConfig creates a new session recording config or overwrites the existing session recording.
UpsertSessionRecordingConfig(ctx context.Context, cfg types.SessionRecordingConfig) (types.SessionRecordingConfig, error)
// DeleteSessionRecordingConfig deletes SessionRecordingConfig from the backend.
DeleteSessionRecordingConfig(ctx context.Context) error
// GetClusterAuditConfig gets ClusterAuditConfig from the backend.
GetClusterAuditConfig(context.Context) (types.ClusterAuditConfig, error)
// CreateClusterAuditConfig creates a cluster audit config if once does not already exist.
CreateClusterAuditConfig(ctx context.Context, cfg types.ClusterAuditConfig) (types.ClusterAuditConfig, error)
// UpdateClusterAuditConfig updates an existing cluster audit config.
UpdateClusterAuditConfig(ctx context.Context, cfg types.ClusterAuditConfig) (types.ClusterAuditConfig, error)
// UpsertClusterAuditConfig creates a new cluster audit config or overwrites the existing cluster audit config.
UpsertClusterAuditConfig(ctx context.Context, cfg types.ClusterAuditConfig) (types.ClusterAuditConfig, error)
// SetClusterAuditConfig sets ClusterAuditConfig from the backend.
SetClusterAuditConfig(context.Context, types.ClusterAuditConfig) error
// DeleteClusterAuditConfig deletes ClusterAuditConfig from the backend.
DeleteClusterAuditConfig(ctx context.Context) error
// GetClusterNetworkingConfig gets ClusterNetworkingConfig from the backend.
GetClusterNetworkingConfig(context.Context) (types.ClusterNetworkingConfig, error)
// CreateClusterNetworkingConfig creates a cluster networking config if once does not already exist.
CreateClusterNetworkingConfig(ctx context.Context, cfg types.ClusterNetworkingConfig) (types.ClusterNetworkingConfig, error)
// UpdateClusterNetworkingConfig updates an existing cluster networking config.
UpdateClusterNetworkingConfig(ctx context.Context, cfg types.ClusterNetworkingConfig) (types.ClusterNetworkingConfig, error)
// UpsertClusterNetworkingConfig creates a new cluster networking config or overwrites the existing cluster networking config.
UpsertClusterNetworkingConfig(ctx context.Context, cfg types.ClusterNetworkingConfig) (types.ClusterNetworkingConfig, error)
// DeleteClusterNetworkingConfig deletes ClusterNetworkingConfig from the backend.
DeleteClusterNetworkingConfig(ctx context.Context) error
// GetInstallers gets all installer scripts from the backend
GetInstallers(context.Context) ([]types.Installer, error)
// ListInstallers returns a page of installer script resources.
ListInstallers(ctx context.Context, limit int, start string) ([]types.Installer, string, error)
// RangeInstallers returns installer script resources within the range [start, end).
RangeInstallers(ctx context.Context, start, end string) iter.Seq2[types.Installer, error]
// GetInstaller gets the installer script from the backend
GetInstaller(ctx context.Context, name string) (types.Installer, error)
// SetInstaller sets the installer script in the backend
SetInstaller(context.Context, types.Installer) error
// DeleteInstaller removes the installer script from the backend
DeleteInstaller(ctx context.Context, name string) error
// DeleteAllInstallers removes all installer script resources from the backend
DeleteAllInstallers(context.Context) error
// GetClusterMaintenanceConfig loads the current maintenance config singleton.
GetClusterMaintenanceConfig(ctx context.Context) (types.ClusterMaintenanceConfig, error)
// UpdateClusterMaintenanceConfig updates the maintenance config singleton.
UpdateClusterMaintenanceConfig(ctx context.Context, cfg types.ClusterMaintenanceConfig) error
// DeleteClusterMaintenanceConfig deletes the maintenance config singleton.
DeleteClusterMaintenanceConfig(ctx context.Context) error
// GetAccessGraphSettings gets the access graph settings from the backend.
GetAccessGraphSettings(context.Context) (*clusterconfigpb.AccessGraphSettings, error)
// CreateAccessGraphSettings creates the access graph settings in the backend.
CreateAccessGraphSettings(context.Context, *clusterconfigpb.AccessGraphSettings) (*clusterconfigpb.AccessGraphSettings, error)
// UpdateAccessGraphSettings updates the access graph settings in the backend.
UpdateAccessGraphSettings(context.Context, *clusterconfigpb.AccessGraphSettings) (*clusterconfigpb.AccessGraphSettings, error)
// UpsertAccessGraphSettings creates or updates the access graph settings in the backend.
UpsertAccessGraphSettings(context.Context, *clusterconfigpb.AccessGraphSettings) (*clusterconfigpb.AccessGraphSettings, error)
// DeleteAccessGraphSettings deletes the access graph settings from the backend.
DeleteAccessGraphSettings(context.Context) error
}
// ClusterConfigurationInternal extends [ClusterConfiguration] with
// auth-specific methods.
type ClusterConfigurationInternal interface {
ClusterConfiguration
StaticScopedTokenService
// SetClusterName sets services.ClusterName on the backend.
SetClusterName(types.ClusterName) error
// UpsertClusterName upserts cluster name
UpsertClusterName(types.ClusterName) error
// DeleteClusterName deletes cluster name resource
DeleteClusterName() error
// AppendCheckAuthPreferenceActions appends some atomic write actions to the
// given slice that will check that the currently stored cluster auth
// preference has the given revision when applied as part of a
// [backend.Backend.AtomicWrite]. The backend to which the actions are
// applied should be the same backend used by the
// ClusterConfigurationInternal.
AppendCheckAuthPreferenceActions(actions []backend.ConditionalAction, revision string) ([]backend.ConditionalAction, error)
}
// StaticScopedTokenService is the interface for interacting with the cluster's
// [*joiningv1.StaticScopedTokens].
type StaticScopedTokenService interface {
// GetStaticScopedTokens gets [*joiningv1.StaticScopedTokens] from the backend.
GetStaticScopedTokens(context.Context) (*joiningv1.StaticScopedTokens, error)
// SetStaticScopedTokens sets [*joiningv1.StaticScopedTokens] to the backend.
SetStaticScopedTokens(context.Context, *joiningv1.StaticScopedTokens) error
// DeleteStaticScopedTokens deletes the [*joiningv1.StaticScopedTokens] resource resource
// from the backend
DeleteStaticScopedTokens(context.Context) error
}
// ValidateAuthPreference performs checks that should happen before persisting a
// new version of the preference resource, typically only as part of Auth
// service operations.
func ValidateAuthPreference(ap types.AuthPreference) error {
// TODO(espadolini): the checks that are duplicated in
// {Set,Create,Update,Upsert}AuthPreference should be moved here
if err := modules.ValidateResource(ap); err != nil {
return trace.Wrap(err)
}
if err := ValidateStableUNIXUserConfig(ap.GetStableUNIXUserConfig()); err != nil {
return trace.Wrap(err)
}
return nil
}
// ValidateStableUNIXUserConfig checks if the configuration is suitable for
// storage and use.
func ValidateStableUNIXUserConfig(c *types.StableUNIXUserConfig) error {
if c == nil || !c.Enabled {
return nil
}
if c.FirstUid > c.LastUid {
return trace.BadParameter("stable UNIX user is enabled but UID range is empty")
}
// see https://github.com/systemd/systemd/blob/cc7300fc5868f6d47f3f47076100b574bf54e58d/docs/UIDS-GIDS.md
const firstUserUID = 1000
if c.FirstUid < firstUserUID {
return trace.BadParameter("stable UNIX user UID range includes negative or system UIDs; the configured range should be contained between 1000 and 2147483647")
}
return nil
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// ConnectionsDiagnostic defines an interface for managing Connection Diagnostics.
type ConnectionsDiagnostic interface {
// CreateConnectionDiagnostic creates a new Connection Diagnostic
CreateConnectionDiagnostic(context.Context, types.ConnectionDiagnostic) error
// UpdateConnectionDiagnostic updates a Connection Diagnostic
UpdateConnectionDiagnostic(context.Context, types.ConnectionDiagnostic) error
// GetConnectionDiagnostic receives a name and returns the Connection Diagnostic matching that name
//
// If not found, a `trace.NotFound` error is returned
GetConnectionDiagnostic(ctx context.Context, name string) (types.ConnectionDiagnostic, error)
// ConnectionDiagnosticTraceAppender adds a method to append traces into ConnectionDiagnostics.
ConnectionDiagnosticTraceAppender
}
// ConnectionDiagnosticTraceAppender specifies methods to add Traces into a DiagnosticConnection
type ConnectionDiagnosticTraceAppender interface {
// AppendDiagnosticTrace atomically adds a new trace into the ConnectionDiagnostic.
AppendDiagnosticTrace(ctx context.Context, name string, t *types.ConnectionDiagnosticTrace) (types.ConnectionDiagnostic, error)
}
// MarshalConnectionDiagnostic marshals the ConnectionDiagnostic resource to JSON.
func MarshalConnectionDiagnostic(s types.ConnectionDiagnostic, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch s := s.(type) {
case *types.ConnectionDiagnosticV1:
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, s))
}
return nil, trace.BadParameter("unrecognized connection diagnostic version %T", s)
}
// UnmarshalConnectionDiagnostic unmarshals the ConnectionDiagnostic resource from JSON.
func UnmarshalConnectionDiagnostic(data []byte, opts ...MarshalOption) (types.ConnectionDiagnostic, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing connection diagnostic data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var s types.ConnectionDiagnosticV1
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("connection diagnostic resource version %q is not supported", h.Version)
}
/*
* Teleport
* Copyright (C) 2024 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"
crownjewelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/crownjewel/v1"
)
// CrownJewels is the interface for managing crown jewel resources.
type CrownJewels interface {
// ListCrownJewels returns the crown jewel resources.
ListCrownJewels(ctx context.Context, pageSize int64, nextToken string) ([]*crownjewelv1.CrownJewel, string, error)
// GetCrownJewel returns the crown jewel resource by name.
GetCrownJewel(ctx context.Context, name string) (*crownjewelv1.CrownJewel, error)
// CreateCrownJewel creates a new crown jewel resource.
CreateCrownJewel(context.Context, *crownjewelv1.CrownJewel) (*crownjewelv1.CrownJewel, error)
// UpdateCrownJewel updates the crown jewel resource.
UpdateCrownJewel(context.Context, *crownjewelv1.CrownJewel) (*crownjewelv1.CrownJewel, error)
// UpsertCrownJewel creates or updates the crown jewel resource.
UpsertCrownJewel(context.Context, *crownjewelv1.CrownJewel) (*crownjewelv1.CrownJewel, error)
// DeleteCrownJewel deletes the crown jewel resource by name.
DeleteCrownJewel(context.Context, string) error
}
// MarshalCrownJewel marshals the CrownJewel object into a JSON byte array.
func MarshalCrownJewel(object *crownjewelv1.CrownJewel, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalCrownJewel unmarshals the CrownJewel object from a JSON byte array.
func UnmarshalCrownJewel(data []byte, opts ...MarshalOption) (*crownjewelv1.CrownJewel, error) {
return UnmarshalProtoResource[*crownjewelv1.CrownJewel](data, opts...)
}
// Teleport
// Copyright (C) 2026 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 (
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend"
)
// GetCursorForResource returns the pagination cursor for a
// resource used in ListResources.
func GetCursorForResource(r types.ResourceWithLabels) string {
switch res := r.(type) {
case types.AppServer:
return GetCursorForAppServer(res)
case types.KubeServer:
return GetCursorForKubeServer(res)
case types.Server:
return GetCursorForNode(res)
}
return backend.GetPaginationKey(r)
}
/*
* 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"
"errors"
"iter"
"log/slog"
"net"
"net/netip"
"net/url"
"slices"
"strings"
"github.com/gravitational/trace"
"go.mongodb.org/mongo-driver/mongo/readpref"
"go.mongodb.org/mongo-driver/x/mongo/driver/connstring"
"github.com/gravitational/teleport/api/types"
azureutils "github.com/gravitational/teleport/api/utils/azure"
gcputils "github.com/gravitational/teleport/api/utils/gcp"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/srv/db/common/enterprise"
"github.com/gravitational/teleport/lib/srv/db/redis/connection"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
awsutils "github.com/gravitational/teleport/lib/utils/aws"
)
// DatabaseGetter defines interface for fetching database resources.
type DatabaseGetter interface {
// GetDatabases returns all database resources.
// Deprecated: Prefer paginated variant such as [ListDatabases] or [RangeDatabases]
GetDatabases(context.Context) ([]types.Database, error)
// ListDatabases returns a page of database resources.
ListDatabases(ctx context.Context, limit int, startKey string) ([]types.Database, string, error)
// RangeDatabases returns database resources within the range [start, end).
RangeDatabases(ctx context.Context, start, end string) iter.Seq2[types.Database, error]
// GetDatabase returns the specified database resource.
GetDatabase(ctx context.Context, name string) (types.Database, error)
}
// Databases defines an interface for managing database resources.
type Databases interface {
// DatabaseGetter provides methods for fetching database resources.
DatabaseGetter
// CreateDatabase creates a new database resource.
CreateDatabase(context.Context, types.Database) error
// UpdateDatabase updates an existing database resource.
UpdateDatabase(context.Context, types.Database) error
// DeleteDatabase removes the specified database resource.
DeleteDatabase(ctx context.Context, name string) error
// DeleteAllDatabases removes all database resources.
DeleteAllDatabases(context.Context) error
}
// MarshalDatabase marshals the database resource to JSON.
func MarshalDatabase(database types.Database, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch database := database.(type) {
case *types.DatabaseV3:
if err := database.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, database))
default:
return nil, trace.BadParameter("unsupported database resource %T", database)
}
}
// UnmarshalDatabase unmarshals the database resource from JSON.
func UnmarshalDatabase(data []byte, opts ...MarshalOption) (types.Database, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing database resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var database types.DatabaseV3
if err := utils.FastUnmarshal(data, &database); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := database.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
database.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
database.SetExpiry(cfg.Expires)
}
return &database, nil
}
return nil, trace.BadParameter("unsupported database resource version %q", h.Version)
}
// ValidateDatabase validates a types.Database.
func ValidateDatabase(db types.Database) error {
if err := enterprise.ProtocolValidation(db.GetProtocol()); err != nil {
return trace.Wrap(err)
}
if err := CheckAndSetDefaults(db); err != nil {
return trace.Wrap(err)
}
if !slices.Contains(defaults.DatabaseProtocols, db.GetProtocol()) {
return trace.BadParameter("unsupported database %q protocol %q, supported are: %v", db.GetName(), db.GetProtocol(), defaults.DatabaseProtocols)
}
// For MongoDB we support specifying either server address or connection
// string in the URI which is useful when connecting to a replica set.
if db.GetProtocol() == defaults.ProtocolMongoDB &&
(strings.HasPrefix(db.GetURI(), connstring.SchemeMongoDB+"://") ||
strings.HasPrefix(db.GetURI(), connstring.SchemeMongoDBSRV+"://")) {
if err := validateMongoDB(db); err != nil {
return trace.Wrap(err)
}
} else if db.GetProtocol() == defaults.ProtocolOracle {
if err := validateOracleURI(db.GetURI()); err != nil {
return trace.BadParameter("invalid Oracle database %q address: %q, error: %v", db.GetName(), db.GetURI(), err)
}
} else if db.GetProtocol() == defaults.ProtocolRedis {
_, err := connection.ParseRedisAddress(db.GetURI())
if err != nil {
return trace.BadParameter("invalid Redis database %q address: %q, error: %v", db.GetName(), db.GetURI(), err)
}
} else if db.GetProtocol() == defaults.ProtocolSnowflake {
if !strings.Contains(db.GetURI(), defaults.SnowflakeURL) {
return trace.BadParameter("Snowflake address should contain " + defaults.SnowflakeURL)
}
} else if db.GetProtocol() == defaults.ProtocolClickHouse || db.GetProtocol() == defaults.ProtocolClickHouseHTTP {
if err := validateClickhouseURI(db); err != nil {
return trace.Wrap(err)
}
} else if db.GetProtocol() == defaults.ProtocolSQLServer {
if err := ValidateSQLServerURI(db.GetURI()); err != nil {
return trace.BadParameter("invalid SQL Server address: %v", err)
}
} else if db.GetProtocol() == defaults.ProtocolSpanner {
if !gcputils.IsSpannerEndpoint(db.GetURI()) {
return trace.BadParameter("GCP Spanner database %q address should be %q",
db.GetName(), gcputils.SpannerEndpoint)
}
} else if db.GetType() == types.DatabaseTypeAlloyDB {
_, err := gcputils.ParseAlloyDBConnectionURI(db.GetURI())
if err != nil {
return trace.Wrap(err, "invalid AlloyDB address")
}
} else if needsURIValidation(db) {
if _, _, err := net.SplitHostPort(db.GetURI()); err != nil {
return trace.BadParameter("invalid database %q address %q: %v", db.GetName(), db.GetURI(), err)
}
}
if db.GetTLS().CACert != "" {
if _, err := tlsca.ParseCertificatePEM([]byte(db.GetTLS().CACert)); err != nil {
return trace.BadParameter("provided database %q CA doesn't appear to be a valid x509 certificate: %v", db.GetName(), err)
}
}
// Validate Active Directory specific configuration, when Kerberos auth is required.
if needsADValidation(db) {
if db.GetAD().KeytabFile == "" && db.GetAD().KDCHostName == "" {
return trace.BadParameter("either keytab file path or kdc_host_name must be provided for database %q, both are missing", db.GetName())
}
if db.GetAD().Krb5File == "" {
return trace.BadParameter("missing Kerberos config file path for database %q", db.GetName())
}
if db.GetAD().Domain == "" {
return trace.BadParameter("missing Active Directory domain for database %q", db.GetName())
}
if db.GetAD().SPN == "" {
return trace.BadParameter("missing service principal name for database %q", db.GetName())
}
if db.GetAD().KDCHostName != "" {
if db.GetAD().LDAPCert == "" {
return trace.BadParameter("missing LDAP certificate for x509 authentication for database %q", db.GetName())
}
if _, err := tlsca.ParseCertificatePEM([]byte(db.GetAD().LDAPCert)); err != nil {
return trace.BadParameter("provided database %q LDAP certificate doesn't appear to be valid: %v", db.GetName(), err)
}
}
}
awsMeta := db.GetAWS()
if awsMeta.AssumeRoleARN != "" {
if awsMeta.AccountID == "" {
return trace.BadParameter("database %q missing AWS account ID", db.GetName())
}
parsed, err := awsutils.ParseRoleARN(awsMeta.AssumeRoleARN)
if err != nil {
return trace.BadParameter("database %q assume_role_arn %q is invalid: %v",
db.GetName(), awsMeta.AssumeRoleARN, err)
}
err = awsutils.CheckARNPartitionAndAccount(parsed, awsMeta.Partition(), awsMeta.AccountID)
if err != nil {
return trace.BadParameter("database %q is incompatible with AWS assume_role_arn %q: %v",
db.GetName(), awsMeta.AssumeRoleARN, err)
}
}
return nil
}
// needsADValidation returns whether a database AD configuration needs to
// be validated.
// We support Azure AD authentication and Kerberos auth with AD for SQL
// Server. The first method doesn't require additional configuration since
// it assumes the environment’s Azure credentials
// (https://learn.microsoft.com/en-us/azure/developer/go/azure-sdk-authentication).
// AD configurations are only required for the second method.
func needsADValidation(db types.Database) bool {
if db.GetProtocol() != defaults.ProtocolSQLServer {
return false
}
// Domain is always required when configuring the AD section, so we assume
// users intend to use Kerberos authentication if the configuration has it.
if db.GetAD().Domain != "" {
return true
}
// Azure-hosted databases and RDS Proxy support other authentication
// methods, and do not require this section to be validated.
if strings.Contains(db.GetURI(), azureutils.MSSQLEndpointSuffix) || db.GetAWS().RDSProxy.Name != "" {
return false
}
return true
}
func validateClickhouseURI(db types.Database) error {
u, err := url.Parse(db.GetURI())
if err != nil {
return trace.BadParameter("failed to parse uri: %v", err)
}
var requiredSchema string
if db.GetProtocol() == defaults.ProtocolClickHouse {
requiredSchema = "clickhouse"
}
if db.GetProtocol() == defaults.ProtocolClickHouseHTTP {
requiredSchema = "https"
}
if u.Scheme != requiredSchema {
return trace.BadParameter("invalid uri schema: %s for %v database protocol", u.Scheme, db.GetProtocol())
}
return nil
}
// needsURIValidation returns whether a database URI needs to be validated.
func needsURIValidation(db types.Database) bool {
switch db.GetProtocol() {
case defaults.ProtocolCassandra, defaults.ProtocolDynamoDB:
// cloud hosted Cassandra doesn't require URI validation.
return db.GetAWS().Region == "" || db.GetAWS().AccountID == ""
default:
return true
}
}
// validateMongoDB validates MongoDB URIs with "mongodb" schemes.
func validateMongoDB(db types.Database) error {
connString, err := connstring.ParseAndValidate(db.GetURI())
// connstring.ParseAndValidate requires DNS resolution on TXT/SRV records
// for a full validation for "mongodb+srv" URIs. We will try to skip the
// DNS errors here by replacing the scheme and then ParseAndValidate again
// to validate as much as we can.
if isDNSError(err) {
slog.WarnContext(context.Background(), "MongoDB database %q (connection string %q) failed validation with DNS error",
"database_name", db.GetName(),
"database_uri", db.GetURI(),
"error", err,
)
connString, err = connstring.ParseAndValidate(strings.Replace(
db.GetURI(),
connstring.SchemeMongoDBSRV+"://",
connstring.SchemeMongoDB+"://",
1,
))
}
if err != nil {
return trace.BadParameter("invalid MongoDB database %q connection string %q: %v", db.GetName(), db.GetURI(), err)
}
// Validate read preference to catch typos early.
if connString.ReadPreference != "" {
if _, err := readpref.ModeFromString(connString.ReadPreference); err != nil {
return trace.BadParameter("invalid MongoDB database %q read preference %q", db.GetName(), connString.ReadPreference)
}
}
return nil
}
// ValidateSQLServerURI validates SQL Server URI and returns host and
// port.
//
// Since Teleport only supports SQL Server authentcation using AD (self-hosted
// or Azure) the database URI must include: computer name, domain and port.
//
// A few examples of valid URIs:
// - computer.ad.example.com:1433
// - computer.domain.com:1433
func ValidateSQLServerURI(uri string) error {
// sqlServerSchema is the SQL Server schema.
const sqlServerSchema = "mssql"
// Add a temporary schema to make a valid URL for url.Parse if schema is
// not found.
if !strings.Contains(uri, "://") {
uri = sqlServerSchema + "://" + uri
}
parsedURI, err := url.Parse(uri)
if err != nil {
return trace.BadParameter("unable to parse database address: %s", err)
}
if parsedURI.Scheme != sqlServerSchema {
return trace.BadParameter("only %q is supported as database address schema", sqlServerSchema)
}
if parsedURI.Port() == "" {
return trace.BadParameter("database address must include port")
}
if parsedURI.Path != "" {
return trace.BadParameter("database address with database name is not supported")
}
if _, err := netip.ParseAddr(parsedURI.Hostname()); err == nil {
return trace.BadParameter("database address as IP is not supported, use URI with domain and computer name instead")
}
parts := strings.Split(parsedURI.Hostname(), ".")
if len(parts) < 3 {
return trace.BadParameter("database address must include domain and computer name")
}
return nil
}
func validateOracleURI(uri string) error {
parts := strings.Split(uri, ",")
for _, part := range parts {
if strings.TrimSpace(part) == "" {
return trace.BadParameter("invalid empty part of URI %q", uri)
}
_, _, err := net.SplitHostPort(part)
if err != nil {
return trace.Wrap(err)
}
}
return nil
}
func isDNSError(err error) bool {
if err == nil {
return false
}
var dnsErr *net.DNSError
return errors.As(err, &dnsErr)
}
const (
// RDSDescribeTypeInstance is used to filter for Instances type of RDS DBs when describing RDS Databases.
RDSDescribeTypeInstance = "instance"
// RDSDescribeTypeCluster is used to filter for Clusters type of RDS DBs when describing RDS Databases.
RDSDescribeTypeCluster = "cluster"
)
const (
// RDSEngineMySQL is RDS engine name for MySQL instances.
RDSEngineMySQL = "mysql"
// RDSEnginePostgres is RDS engine name for Postgres instances.
RDSEnginePostgres = "postgres"
// RDSEngineMariaDB is RDS engine name for MariaDB instances.
RDSEngineMariaDB = "mariadb"
// RDSEngineAurora is RDS engine name for Aurora MySQL 5.6 compatible clusters.
// This reached EOF on Feb 28, 2023.
// https://docs.aws.amazon.com/AmazonRDS/latest/AuroraUserGuide/Aurora.MySQL56.EOL.html
RDSEngineAurora = "aurora"
// RDSEngineAuroraMySQL is RDS engine name for Aurora MySQL 5.7 compatible clusters.
RDSEngineAuroraMySQL = "aurora-mysql"
// RDSEngineAuroraPostgres is RDS engine name for Aurora Postgres clusters.
RDSEngineAuroraPostgres = "aurora-postgresql"
)
const (
// RDSEngineModeProvisioned is the RDS engine mode for provisioned Aurora clusters
RDSEngineModeProvisioned = "provisioned"
// RDSEngineModeServerless is the RDS engine mode for Aurora Serverless DB clusters
RDSEngineModeServerless = "serverless"
// RDSEngineModeParallelQuery is the RDS engine mode for Aurora MySQL clusters with parallel query enabled
RDSEngineModeParallelQuery = "parallelquery"
)
const (
// RDSProxyMySQLPort is the port that RDS Proxy listens on for MySQL connections.
RDSProxyMySQLPort = 3306
// RDSProxyPostgresPort is the port that RDS Proxy listens on for Postgres connections.
RDSProxyPostgresPort = 5432
// RDSProxySQLServerPort is the port that RDS Proxy listens on for SQL Server connections.
RDSProxySQLServerPort = 1433
)
const (
// AzureEngineMySQL is the Azure engine name for MySQL single-server instances.
AzureEngineMySQL = "Microsoft.DBforMySQL/servers"
// AzureEngineMySQLFlex is the Azure engine name for MySQL flexible-server instances.
AzureEngineMySQLFlex = "Microsoft.DBforMySQL/flexibleServers"
// AzureEnginePostgres is the Azure engine name for PostgreSQL single-server instances.
AzureEnginePostgres = "Microsoft.DBforPostgreSQL/servers"
// AzureEnginePostgresFlex is the Azure engine name for PostgreSQL flexible-server instances.
AzureEnginePostgresFlex = "Microsoft.DBforPostgreSQL/flexibleServers"
)
const (
// RedshiftServerlessWorkgroupEndpoint is the endpoint type for workgroups.
RedshiftServerlessWorkgroupEndpoint = "workgroup"
// RedshiftServerlessVPCEndpoint is the endpoint type for VCP endpoints.
RedshiftServerlessVPCEndpoint = "vpc-endpoint"
)
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// MarshalDatabaseServer marshals the DatabaseServer resource to JSON.
func MarshalDatabaseServer(databaseServer types.DatabaseServer, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch databaseServer := databaseServer.(type) {
case *types.DatabaseServerV3:
if err := databaseServer.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, databaseServer))
default:
return nil, trace.BadParameter("unrecognized database server version %T", databaseServer)
}
}
// UnmarshalDatabaseServer unmarshals the DatabaseServer resource from JSON.
func UnmarshalDatabaseServer(data []byte, opts ...MarshalOption) (types.DatabaseServer, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing database server data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var s types.DatabaseServerV3
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("database server resource version %q is not supported", h.Version)
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// DatabaseServices defines an interface for managing DatabaseService resources.
type DatabaseServices interface {
// UpsertDatabaseService updates an existing DatabaseService resource.
UpsertDatabaseService(context.Context, types.DatabaseService) (*types.KeepAlive, error)
// DeleteDatabaseService removes the specified DatabaseService resource.
DeleteDatabaseService(ctx context.Context, name string) error
// DeleteAllDatabaseServices removes all DatabaseService resources.
DeleteAllDatabaseServices(context.Context) error
}
// MarshalDatabaseService marshals the DatabaseService resource to JSON.
func MarshalDatabaseService(databaseService types.DatabaseService, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch databaseService := databaseService.(type) {
case *types.DatabaseServiceV1:
if err := databaseService.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, databaseService))
default:
return nil, trace.BadParameter("unrecognized DatabaseService version %T", databaseService)
}
}
// UnmarshalDatabaseService unmarshals the DatabaseService resource from JSON.
func UnmarshalDatabaseService(data []byte, opts ...MarshalOption) (types.DatabaseService, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing DatabaseService data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var s types.DatabaseServiceV1
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("database service resource version %q is not supported", h.Version)
}
// Teleport
// Copyright (C) 2026 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/trace"
delegationv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/delegation/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend"
)
// DelegationSessions is an interface over the DelegationSessions service. This
// interface may also be implemented by a client to allow remote and local
// consumers to access the resource in a similar way.
type DelegationSessions interface {
// CreateDelegationSession creates a new delegation session.
CreateDelegationSession(ctx context.Context, session *delegationv1.DelegationSession) (*delegationv1.DelegationSession, error)
// GetDelegationSession reads a delegation session using its ID.
GetDelegationSession(ctx context.Context, id string) (*delegationv1.DelegationSession, error)
// DeleteDelegationSession deletes a delegation session using its ID.
DeleteDelegationSession(ctx context.Context, id string) error
// AppendPutDelegationSessionActions adds conditional actions to an atomic
// write to create or update a DelegationSession.
AppendPutDelegationSessionActions(
actions []backend.ConditionalAction,
session *delegationv1.DelegationSession,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteDelegationSessionActions adds conditional actions to an atomic
// write to delete a DelegationSession.
AppendDeleteDelegationSessionActions(
actions []backend.ConditionalAction,
id string,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
// ValidateDelegationSession validates a DelegationSession object.
func ValidateDelegationSession(p *delegationv1.DelegationSession) error {
switch {
case p == nil:
return trace.BadParameter("must not be nil")
case p.GetKind() != types.KindDelegationSession:
return trace.BadParameter("kind: must be %s", types.KindDelegationSession)
case p.GetVersion() != types.V1:
return trace.BadParameter("version: must be %s", types.V1)
case p.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name: is required")
case p.GetMetadata().GetExpires() == nil:
return trace.BadParameter("metadata.expires: is required")
case p.GetSpec().GetUser() == "":
return trace.BadParameter("spec.user: is required")
}
if len(p.GetSpec().GetResources()) == 0 {
return trace.BadParameter("spec.resources: at least one resource is required")
}
var hasWildcard, hasExplicit bool
for idx, spec := range p.GetSpec().GetResources() {
if err := ValidateDelegationResourceSpec(spec); err != nil {
return trace.BadParameter("spec.resources[%d]: invalid resource spec: %v", idx, err)
}
if spec.GetKind() == types.Wildcard {
hasWildcard = true
} else {
hasExplicit = true
}
if hasWildcard && hasExplicit {
return trace.BadParameter("spec.resources: wildcard is mutually exclusive with explicit resources")
}
}
if len(p.GetSpec().GetAuthorizedUsers()) == 0 {
return trace.BadParameter("spec.authorized_users: at least one user is required")
}
for idx, user := range p.GetSpec().GetAuthorizedUsers() {
if user.GetKind() != types.KindBot {
return trace.BadParameter("spec.authorized_users[%d].kind: must be %s", idx, types.KindBot)
}
if user.GetBotName() == "" {
return trace.BadParameter("spec.authorized_users[%d].bot_name: is required", idx)
}
}
return nil
}
// ValidateDelegationResourceSpec validates a DelegationResourceSpec object.
func ValidateDelegationResourceSpec(s *delegationv1.DelegationResourceSpec) error {
if s.GetName() == "" {
return trace.BadParameter("name is required")
}
// TODO(boxofrad): implement support for constraints.
if s.GetConstraints() != nil {
return trace.BadParameter("constraints are not yet supported")
}
switch s.GetKind() {
case types.KindApp, types.KindDatabase, types.KindNode, types.KindKubernetesCluster, types.KindWindowsDesktop, types.KindGitServer, types.Wildcard:
case "":
return trace.BadParameter("kind is required")
default:
return trace.BadParameter("invalid kind: %q", s.GetKind())
}
switch {
case s.GetKind() == types.Wildcard && s.GetName() != types.Wildcard:
return trace.BadParameter("name must also be '*'")
case s.GetKind() != types.Wildcard && s.GetName() == types.Wildcard:
return trace.BadParameter("kind must also be '*'")
}
return nil
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// WindowsDesktops defines an interface for managing Windows desktop hosts.
type WindowsDesktops interface {
WindowsDesktopGetter
CreateWindowsDesktop(context.Context, types.WindowsDesktop) error
UpdateWindowsDesktop(context.Context, types.WindowsDesktop) error
UpsertWindowsDesktop(ctx context.Context, desktop types.WindowsDesktop) error
DeleteWindowsDesktop(ctx context.Context, hostID, name string) error
DeleteAllWindowsDesktops(context.Context) error
ListWindowsDesktops(ctx context.Context, req types.ListWindowsDesktopsRequest) (*types.ListWindowsDesktopsResponse, error)
ListWindowsDesktopServices(ctx context.Context, req types.ListWindowsDesktopServicesRequest) (*types.ListWindowsDesktopServicesResponse, error)
}
// WindowsDesktopGetter is an interface for fetching WindowsDesktop resources.
type WindowsDesktopGetter interface {
GetWindowsDesktops(context.Context, types.WindowsDesktopFilter) ([]types.WindowsDesktop, error)
}
// MarshalWindowsDesktop marshals the WindowsDesktop resource to JSON.
func MarshalWindowsDesktop(s types.WindowsDesktop, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch s := s.(type) {
case *types.WindowsDesktopV3:
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, s))
default:
return nil, trace.BadParameter("unrecognized windows desktop version %T", s)
}
}
// UnmarshalWindowsDesktop unmarshals the WindowsDesktop resource from JSON.
func UnmarshalWindowsDesktop(data []byte, opts ...MarshalOption) (types.WindowsDesktop, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing windows desktop data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var s types.WindowsDesktopV3
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("windows desktop resource version %q is not supported", h.Version)
}
// MarshalWindowsDesktopService marshals the WindowsDesktopService resource to JSON.
func MarshalWindowsDesktopService(s types.WindowsDesktopService, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch s := s.(type) {
case *types.WindowsDesktopServiceV3:
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, s))
default:
return nil, trace.BadParameter("unrecognized windows desktop service version %T", s)
}
}
// UnmarshalWindowsDesktopService unmarshals the WindowsDesktopService resource from JSON.
func UnmarshalWindowsDesktopService(data []byte, opts ...MarshalOption) (types.WindowsDesktopService, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing windows desktop service data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var s types.WindowsDesktopServiceV3
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("windows desktop service resource version %q is not supported", h.Version)
}
/*
* 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/trace"
devicepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/devicetrust/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/utils"
)
// DevicesGetter allows to list all registered devices from storage.
type DevicesGetter interface {
ListDevices(ctx context.Context, pageSize int, pageToken string, view devicepb.DeviceView) (devices []*devicepb.Device, nextPageToken string, err error)
}
// UnmarshalDevice unmarshals a DeviceV1 resource and runs CheckAndSetDefaults.
func UnmarshalDevice(raw []byte) (*types.DeviceV1, error) {
dev := &types.DeviceV1{}
if err := utils.FastUnmarshal(raw, dev); err != nil {
return nil, trace.Wrap(err)
}
return dev, trace.Wrap(dev.CheckAndSetDefaults())
}
// MarshalDevice marshals a DeviceV1 resource.
func MarshalDevice(dev *types.DeviceV1) ([]byte, error) {
devBytes, err := utils.FastMarshal(dev)
if err != nil {
return nil, trace.Wrap(err)
}
return devBytes, nil
}
var (
// unmarshalDeviceFromBackendItemConv is a convenience function that converts
// a backend.Item to a *devicepb.Device.
// It's populated when e/lib/devicetrust/storage/ is initialized.
unmarshalDeviceFromBackendItemConv func(item backend.Item) (*devicepb.Device, error)
)
// SetUnmarshalDeviceFromBackendItemConv allows to set a custom conversion function for
// unmarshaling a [devicepb.Device] from a [backend.Item].
// This function must be called in the init() function of the package that owns the conversion logic.
// It's not safe to call this function concurrently.
func SetUnmarshalDeviceFromBackendItemConv(conv func(item backend.Item) (*devicepb.Device, error)) {
unmarshalDeviceFromBackendItemConv = conv
}
// UnmarshalDeviceFromBackendItem unmarshals a devicepb.Device from a backend.Item.
// It's a convenience function that uses UnmarshalDeviceFromBackendItemConv because
// the storage package uses an internal representation of devicepb.Device when storing
// it in the backend.
func UnmarshalDeviceFromBackendItem(item backend.Item) (*devicepb.Device, error) {
if unmarshalDeviceFromBackendItemConv == nil {
return nil, trace.BadParameter("UnmarshalDeviceFromBackendItemConv is not set")
}
res, err := unmarshalDeviceFromBackendItemConv(item)
return res, trace.Wrap(err)
}
/*
* 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/trace"
discoveryconfigclient "github.com/gravitational/teleport/api/client/discoveryconfig"
"github.com/gravitational/teleport/api/types/discoveryconfig"
"github.com/gravitational/teleport/lib/utils"
)
var _ DiscoveryConfigs = (*discoveryconfigclient.Client)(nil)
// DiscoveryConfigs defines an interface for managing DiscoveryConfigs.
type DiscoveryConfigs interface {
DiscoveryConfigsGetter
// CreateDiscoveryConfig creates a new DiscoveryConfig resource.
CreateDiscoveryConfig(context.Context, *discoveryconfig.DiscoveryConfig) (*discoveryconfig.DiscoveryConfig, error)
// UpdateDiscoveryConfig updates an existing DiscoveryConfig resource.
UpdateDiscoveryConfig(context.Context, *discoveryconfig.DiscoveryConfig) (*discoveryconfig.DiscoveryConfig, error)
// UpsertDiscoveryConfig upserts a DiscoveryConfig resource.
UpsertDiscoveryConfig(context.Context, *discoveryconfig.DiscoveryConfig) (*discoveryconfig.DiscoveryConfig, error)
// DeleteDiscoveryConfig removes the specified DiscoveryConfig resource.
DeleteDiscoveryConfig(ctx context.Context, name string) error
// DeleteAllDiscoveryConfigs removes all DiscoveryConfigs.
DeleteAllDiscoveryConfigs(context.Context) error
}
// DiscoveryConfigWithStatusUpdater defines an interface for managing DiscoveryConfig resources including updating their status.
type DiscoveryConfigWithStatusUpdater interface {
DiscoveryConfigs
// UpdateDiscoveryConfigStatus updates the status of the specified DiscoveryConfig resource.
UpdateDiscoveryConfigStatus(context.Context, string, discoveryconfig.Status) (*discoveryconfig.DiscoveryConfig, error)
}
// DiscoveryConfigsGetter defines methods for List/Read operations on DiscoveryConfig Resources.
type DiscoveryConfigsGetter interface {
// ListDiscoveryConfigs returns a paginated list of all DiscoveryConfig resources.
// An optional DiscoveryGroup can be provided to filter.
ListDiscoveryConfigs(ctx context.Context, pageSize int, nextToken string) ([]*discoveryconfig.DiscoveryConfig, string, error)
// GetDiscoveryConfig returns the specified DiscoveryConfig resources.
GetDiscoveryConfig(ctx context.Context, name string) (*discoveryconfig.DiscoveryConfig, error)
}
// MarshalDiscoveryConfig marshals the DiscoveryConfig resource to JSON.
func MarshalDiscoveryConfig(discoveryConfig *discoveryconfig.DiscoveryConfig, opts ...MarshalOption) ([]byte, error) {
if err := discoveryConfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *discoveryConfig
copy.SetRevision("")
discoveryConfig = ©
}
return utils.FastMarshal(discoveryConfig)
}
// UnmarshalDiscoveryConfig unmarshals the DiscoveryConfig resource from JSON.
func UnmarshalDiscoveryConfig(data []byte, opts ...MarshalOption) (*discoveryconfig.DiscoveryConfig, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing discovery config data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var discoveryConfig *discoveryconfig.DiscoveryConfig
if err := utils.FastUnmarshal(data, &discoveryConfig); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := discoveryConfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
discoveryConfig.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
discoveryConfig.SetExpiry(cfg.Expires)
}
return discoveryConfig, nil
}
/**
* Teleport
* Copyright (C) 2024 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// DynamicWindowsDesktops defines an interface for managing dynamic Windows desktops.
type DynamicWindowsDesktops interface {
GetDynamicWindowsDesktop(ctx context.Context, name string) (types.DynamicWindowsDesktop, error)
CreateDynamicWindowsDesktop(context.Context, types.DynamicWindowsDesktop) (types.DynamicWindowsDesktop, error)
UpdateDynamicWindowsDesktop(context.Context, types.DynamicWindowsDesktop) (types.DynamicWindowsDesktop, error)
UpsertDynamicWindowsDesktop(context.Context, types.DynamicWindowsDesktop) (types.DynamicWindowsDesktop, error)
DeleteDynamicWindowsDesktop(ctx context.Context, name string) error
ListDynamicWindowsDesktops(ctx context.Context, pageSize int, pageToken string) ([]types.DynamicWindowsDesktop, string, error)
}
// MarshalDynamicWindowsDesktop marshals the DynamicWindowsDesktop resource to JSON.
func MarshalDynamicWindowsDesktop(s types.DynamicWindowsDesktop, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch s := s.(type) {
case *types.DynamicWindowsDesktopV1:
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, s))
default:
return nil, trace.BadParameter("unrecognized windows desktop version %T", s)
}
}
// UnmarshalDynamicWindowsDesktop unmarshals the DynamicWindowsDesktop resource from JSON.
func UnmarshalDynamicWindowsDesktop(data []byte, opts ...MarshalOption) (types.DynamicWindowsDesktop, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing windows desktop data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var s types.DynamicWindowsDesktopV1
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("windows desktop resource version %q is not supported", h.Version)
}
// Teleport
// Copyright (C) 2026 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"
devicepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/devicetrust/v1"
)
// EnrollPairing manages mobile device enrollment pairings.
type EnrollPairing interface {
// CreateEnrollPairing creates a new EnrollPairing for user in the
// AWAITING_DEVICE state with a short TTL.
// Returns AlreadyExists if a pairing already exists for user.
CreateEnrollPairing(ctx context.Context, user string) (*devicepb.EnrollPairing, error)
// GetCurrentEnrollPairing returns the EnrollPairing for user.
// Returns NotFound if no pairing exists.
GetCurrentEnrollPairing(ctx context.Context, user string) (*devicepb.EnrollPairing, error)
// GetEnrollPairingByToken returns the EnrollPairing whose status token
// matches token. Returns NotFound if no pairing matches.
GetEnrollPairingByToken(ctx context.Context, token string) (*devicepb.EnrollPairing, error)
// RequestEnrollPairingApproval transitions pairing from AWAITING_DEVICE to
// AWAITING_APPROVAL, persisting device for the Web UI to display and for
// retry gating, and returns the updated pairing.
// Returns CompareFailed if the pairing is no longer awaiting a device.
RequestEnrollPairingApproval(ctx context.Context, pairing *devicepb.EnrollPairing, device *devicepb.EnrollPairingDevice) (*devicepb.EnrollPairing, error)
// ApproveEnrollPairing transitions pairing from AWAITING_APPROVAL to APPROVED
// and returns the updated pairing.
// Returns CompareFailed if the pairing is no longer awaiting approval.
ApproveEnrollPairing(ctx context.Context, pairing *devicepb.EnrollPairing) (*devicepb.EnrollPairing, error)
// DeleteEnrollPairing removes pairing along with its token index. It backs
// both denial by the user and the single-use consumption of a pairing when
// the enrollment token is issued.
// Returns CompareFailed if the pairing has changed since it was read.
DeleteEnrollPairing(ctx context.Context, pairing *devicepb.EnrollPairing) error
}
// MarshalEnrollPairing marshals an EnrollPairing resource to JSON.
func MarshalEnrollPairing(pairing *devicepb.EnrollPairing, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(pairing, opts...)
}
// UnmarshalEnrollPairing unmarshals an EnrollPairing resource from JSON.
func UnmarshalEnrollPairing(data []byte, opts ...MarshalOption) (*devicepb.EnrollPairing, error) {
return UnmarshalProtoResource[*devicepb.EnrollPairing](data, opts...)
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types/externalauditstorage"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalExternalAuditStorage unmarshals the External Audit Storage resource from JSON.
func UnmarshalExternalAuditStorage(data []byte, opts ...MarshalOption) (*externalauditstorage.ExternalAuditStorage, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing External Audit Storage data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var out *externalauditstorage.ExternalAuditStorage
if err := utils.FastUnmarshal(data, &out); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := out.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
out.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
out.SetExpiry(cfg.Expires)
}
return out, nil
}
// MarshalExternalAuditStorage marshals the External Audit Storage resource to JSON.
func MarshalExternalAuditStorage(externalAuditStorage *externalauditstorage.ExternalAuditStorage, opts ...MarshalOption) ([]byte, error) {
if err := externalAuditStorage.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *externalAuditStorage
copy.SetRevision("")
externalAuditStorage = ©
}
return utils.FastMarshal(externalAuditStorage)
}
/*
* 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"
"errors"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/internalutils/stream"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/scopes"
fb "github.com/gravitational/teleport/lib/utils/fanoutbuffer"
)
var errFanoutReset = errors.New("event fanout system reset")
var errFanoutClosed = errors.New("event fanout system closed")
var errWatcherClosed = errors.New("event watcher closed")
type FanoutV2Config struct {
Capacity uint64
GracePeriod time.Duration
Clock clockwork.Clock
}
func (c *FanoutV2Config) SetDefaults() {
if c.Capacity == 0 {
c.Capacity = 1024
}
if c.GracePeriod == 0 {
// the most frequent periodic writes happen once per minute. a grace period of 59s is a
// reasonable default, since a cursor that can't catch up within 59s is likely to continue
// to fall further behind.
c.GracePeriod = 59 * time.Second
}
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
}
// FanoutV2 is a drop-in replacement for Fanout that offers a different set of performance characteristics. It
// supports variable-size buffers to better accommodate large spikes in event load, but it does so at the cost
// of higher levels of context-switching since all readers are notified of all events as well as higher baseline
// memory usage due to relying on a large shared buffer.
type FanoutV2 struct {
cfg FanoutV2Config
rw sync.RWMutex
buf *fb.Buffer[fanoutV2Entry]
init *fanoutV2Init
closed bool
}
// NewFanoutV2 allocates a new fanout instance.
func NewFanoutV2(cfg FanoutV2Config) *FanoutV2 {
cfg.SetDefaults()
f := &FanoutV2{
cfg: cfg,
}
f.setup()
return f
}
// NewStream gets a new event stream. The provided context will form the basis of
// the stream's close context. Note that streams *must* be explicitly closed when
// completed in order to avoid performance issues.
func (f *FanoutV2) NewStream(ctx context.Context, watch types.Watch) stream.Stream[types.Event] {
f.rw.RLock()
defer f.rw.RUnlock()
if f.closed {
return stream.Fail[types.Event](errFanoutClosed)
}
return &fanoutV2Stream{
closeContext: ctx,
cursor: f.buf.NewCursor(),
init: f.init,
watch: watch,
}
}
func (f *FanoutV2) NewWatcher(ctx context.Context, watch types.Watch) (types.Watcher, error) {
ctx, cancel := context.WithCancel(ctx)
w := &streamWatcher{
cancel: cancel,
events: make(chan types.Event, 16),
// note that we don't use ctx.Done() because we want to wait until
// we've finished stream closure and extracted the resulting error
// before signaling watcher closure.
done: make(chan struct{}),
}
go w.run(ctx, f.NewStream(ctx, watch))
return w, nil
}
type streamWatcher struct {
cancel context.CancelFunc
events chan types.Event
done chan struct{}
emux sync.Mutex
err error
}
func (w *streamWatcher) run(ctx context.Context, stream stream.Stream[types.Event]) {
defer func() {
if err := stream.Done(); err != nil {
w.emux.Lock()
w.err = err
w.emux.Unlock()
}
close(w.done)
}()
for stream.Next() {
select {
case w.events <- stream.Item():
case <-ctx.Done():
return
}
}
}
func (w *streamWatcher) Events() <-chan types.Event {
return w.events
}
func (w *streamWatcher) Done() <-chan struct{} {
return w.done
}
func (w *streamWatcher) Close() error {
w.cancel()
return nil
}
func (w *streamWatcher) Error() error {
w.emux.Lock()
defer w.emux.Unlock()
if w.err != nil {
return w.err
}
select {
case <-w.Done():
return errWatcherClosed
default:
return nil
}
}
func (f *FanoutV2) Emit(events ...types.Event) {
f.rw.RLock()
defer f.rw.RUnlock()
if !f.init.isInit() {
panic("Emit called on uninitialized fanout instance")
}
if f.closed {
// emit racing with close is fairly common with how we
// use this type, so its best to ignore it.
return
}
// batch-process events to minimize the need to acquire the
// fanout buffer's write lock (batching writes has a non-trivial
// impact on fanout buffer benchmarks due to each cursor needing
// to acquire the read lock individually).
var ebuf [16]fanoutV2Entry
for len(events) > 0 {
n := min(len(events), len(ebuf))
for i := range n {
ebuf[i] = newFanoutV2Entry(events[i])
}
f.buf.Append(ebuf[:n]...)
events = events[n:]
}
}
func (f *FanoutV2) Reset() {
f.rw.Lock()
defer f.rw.Unlock()
if f.closed {
return
}
f.teardown(errFanoutReset)
f.setup()
}
func (f *FanoutV2) Close() error {
f.rw.Lock()
defer f.rw.Unlock()
if f.closed {
return nil
}
f.teardown(errFanoutClosed)
f.closed = true
return nil
}
func (f *FanoutV2) setup() {
f.init = newFanoutV2Init()
f.buf = fb.NewBuffer[fanoutV2Entry](fb.Config{
Capacity: f.cfg.Capacity,
GracePeriod: f.cfg.GracePeriod,
Clock: f.cfg.Clock,
})
}
func (f *FanoutV2) teardown(err error) {
f.init.setErr(err)
f.buf.Close()
}
func (f *FanoutV2) SetInit(kinds []types.WatchKind) {
f.rw.RLock()
defer f.rw.RUnlock()
km := make(map[resourceKind]types.WatchKind, len(kinds))
for _, kind := range kinds {
km[resourceKind{kind: kind.Kind, subKind: kind.SubKind}] = kind
}
f.init.setInit(km)
}
// fanoutV2Stream is a stream.Stream implementation that streams events from a FanoutV2 instance. It handles filtering
// out events that don't match the provided watch parameters, and construction of custom init events.
type fanoutV2Stream struct {
closeContext context.Context
cursor *fb.Cursor[fanoutV2Entry]
init *fanoutV2Init
watch types.Watch
rbuf [16]fanoutV2Entry
n, next int
event types.Event
err error
}
func (s *fanoutV2Stream) Next() (ok bool) {
if s.init != nil {
s.event, s.err = s.waitInit(s.closeContext)
s.init = nil
return s.err == nil
}
for {
// try finding the next matching event within read buffer
var ok bool
s.event, ok, s.err = s.advance()
if ok {
return true
}
// read a new batch of events into the read buffer
s.next = 0
s.n, s.err = s.cursor.Read(s.closeContext, s.rbuf[:])
if s.err != nil {
if errors.Is(s.err, fb.ErrBufferClosed) {
s.err = errFanoutReset
}
return false
}
}
}
func (s *fanoutV2Stream) Item() types.Event {
return s.event
}
func (s *fanoutV2Stream) Done() error {
s.cursor.Close()
return s.err
}
// waitInit waits for fanout initialization and builds an appropriate init event.
func (s *fanoutV2Stream) waitInit(ctx context.Context) (types.Event, error) {
confirmedKinds, err := s.init.wait(ctx)
if err != nil {
return types.Event{}, trace.Wrap(err)
}
validKinds := make([]types.WatchKind, 0, len(s.watch.Kinds))
for _, requested := range s.watch.Kinds {
k := resourceKind{kind: requested.Kind, subKind: requested.SubKind}
if configured, ok := confirmedKinds[k]; !ok || !configured.Contains(requested) {
if s.watch.AllowPartialSuccess {
continue
}
return types.Event{}, trace.BadParameter("resource type %q is not supported by this event stream", requested.Kind)
}
validKinds = append(validKinds, requested)
}
if len(validKinds) == 0 {
return types.Event{}, trace.BadParameter("none of the requested resources are supported by this fanoutWatcher")
}
return types.Event{Type: types.OpInit, Resource: types.NewWatchStatus(validKinds)}, nil
}
// advance advances through the stream's internal read buffer looking for the
// next event that matches our specific watch parameters.
func (f *fanoutV2Stream) advance() (event types.Event, ok bool, err error) {
for f.next < f.n {
entry := f.rbuf[f.next]
f.next++
if entry.Event.Resource == nil {
// events with no associated resources are special cases (e.g. OpUnreliable), and are
// emitted to all watchers.
return entry.Event, true, nil
}
for _, kind := range f.watch.Kinds {
match, err := kind.Matches(entry.Event)
if err != nil {
return types.Event{}, false, trace.Wrap(err)
}
if !match {
continue
}
// scope filtering is applied here rather than inside kind.Matches because the scope-matching
// logic lives in lib/scopes, which the api/types package that defines Matches cannot import.
if !WatchKindMatchesScope(kind, entry.Event.Resource) {
continue
}
if kind.LoadSecrets {
return entry.EventWithSecrets, true, nil
}
return entry.Event, true, nil
}
}
return types.Event{}, false, nil
}
// WatchKindMatchesScope reports whether the given resource satisfies the watch kind's scope filter. A
// nil scope filter matches everything. The resource's scope is read via the optional GetScope accessor
// (implemented by scoped resources and the Resource153 legacy adapter); resources that are not scope-aware
// report an empty scope, which scopes.MatchScope treats as unscoped.
//
// This is the scope-filtering counterpart to [types.WatchKind.Matches] and must be applied by every
// event source that honors watch kinds (the fanout and the local backend event parsers), so that a
// scope filter selects the same set of events regardless of which source serves the watch.
//
// TODO(fspmarshall/scopes): at some point it may be worth moving core scope-matching logic to api
// so that we can fold this functionality back into [WatchKind.Matches].
func WatchKindMatchesScope(kind types.WatchKind, resource types.Resource) bool {
if kind.ScopeFilter == nil {
return true
}
var scope string
if scoped, ok := resource.(interface{ GetScope() string }); ok {
scope = scoped.GetScope()
}
return scopes.MatchScope(kind.ScopeFilter.ToProto(), scope)
}
// fanoutV2Entry is the underlying buffer entry that is fanned out to all
// cursors. Individual streams decide if they care about the version of the
// event with or without secrets based on their parameters.
type fanoutV2Entry struct {
Event types.Event
EventWithSecrets types.Event
}
func newFanoutV2Entry(event types.Event) fanoutV2Entry {
return fanoutV2Entry{
Event: filterEventSecrets(event),
EventWithSecrets: event,
}
}
func filterEventSecrets(event types.Event) types.Event {
if r, ok := event.Resource.(types.ResourceWithSecrets); ok {
event.Resource = r.WithoutSecrets()
}
// WebSessions do not implement the ResourceWithSecrets interface.
if r, ok := event.Resource.(types.WebSession); ok {
event.Resource = r.WithoutSecrets()
}
return event
}
type resourceKind struct {
kind string
subKind string
}
// fanoutV2Init is a helper for blocking on and distributing the init event for a fanout
// instance. It uses a channel as both the init signal and a memory barrier to ensure
// good concurrent performance, and it is allocated behind a pointer so that it can be
// easily termianted and replaced during resets, ensuring that we don't need to handle
// edge-cases around old streams observing the wrong event/error.
type fanoutV2Init struct {
once sync.Once
ch chan struct{}
kinds map[resourceKind]types.WatchKind
err error
}
func newFanoutV2Init() *fanoutV2Init {
return &fanoutV2Init{
ch: make(chan struct{}),
}
}
func (i *fanoutV2Init) setInit(kinds map[resourceKind]types.WatchKind) {
i.once.Do(func() {
i.kinds = kinds
close(i.ch)
})
}
func (i *fanoutV2Init) setErr(err error) {
i.once.Do(func() {
i.err = err
close(i.ch)
})
}
func (i *fanoutV2Init) wait(ctx context.Context) (kinds map[resourceKind]types.WatchKind, err error) {
select {
case <-i.ch:
return i.kinds, i.err
case <-ctx.Done():
return nil, trace.Wrap(ctx.Err())
}
}
func (i *fanoutV2Init) isInit() bool {
select {
case <-i.ch:
return true
default:
return false
}
}
/*
* Teleport
* Copyright (C) 2024 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/client/gitserver"
"github.com/gravitational/teleport/api/types"
)
// GitServerGetter defines interface for fetching git servers.
type GitServerGetter gitserver.ReadOnlyClient
// GitServers defines an interface for managing git servers.
type GitServers interface {
GitServerGetter
// CreateGitServer creates a Git server resource.
CreateGitServer(ctx context.Context, item types.Server) (types.Server, error)
// UpdateGitServer updates a Git server resource.
UpdateGitServer(ctx context.Context, item types.Server) (types.Server, error)
// UpsertGitServer updates a Git server resource, creating it if it doesn't exist.
UpsertGitServer(ctx context.Context, item types.Server) (types.Server, error)
// DeleteGitServer removes the specified Git server resource.
DeleteGitServer(ctx context.Context, name string) error
}
// MarshalGitServer marshals the Git Server resource to JSON.
func MarshalGitServer(server types.Server, opts ...MarshalOption) ([]byte, error) {
return MarshalServer(server, opts...)
}
// UnmarshalGitServer unmarshals the Git Server resource from JSON.
func UnmarshalGitServer(bytes []byte, opts ...MarshalOption) (types.Server, error) {
return UnmarshalServer(bytes, types.KindGitServer, opts...)
}
/*
* 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 (
"encoding/json"
"fmt"
"sync"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/utils"
)
// ErrRequiresEnterprise indicates that a feature requires
// Teleport Enterprise.
var ErrRequiresEnterprise = &trace.AccessDeniedError{Message: "this feature requires Teleport Enterprise"}
// githubConnectorMutex is a mutex for the GitHub auth connector
// registration functions.
var githubConnectorMutex sync.RWMutex
// GithubAuthCreator creates a new GitHub connector.
type GithubAuthCreator func(string, types.GithubConnectorSpecV3) (types.GithubConnector, error)
var githubAuthCreator GithubAuthCreator
// RegisterGithubAuthCreator registers a function to create GitHub auth connectors.
func RegisterGithubAuthCreator(creator GithubAuthCreator) {
githubConnectorMutex.Lock()
defer githubConnectorMutex.Unlock()
githubAuthCreator = creator
}
// NewGithubConnector creates a new GitHub auth connector.
func NewGithubConnector(name string, spec types.GithubConnectorSpecV3) (types.GithubConnector, error) {
githubConnectorMutex.RLock()
defer githubConnectorMutex.RUnlock()
return githubAuthCreator(name, spec)
}
// GithubAuthInitializer initializes a GitHub auth connector.
type GithubAuthInitializer func(types.GithubConnector) (types.GithubConnector, error)
var githubAuthInitializer GithubAuthInitializer
// RegisterGithubAuthInitializer registers a function to initialize GitHub auth connectors.
func RegisterGithubAuthInitializer(init GithubAuthInitializer) {
githubConnectorMutex.Lock()
defer githubConnectorMutex.Unlock()
githubAuthInitializer = init
}
// InitGithubConnector initializes c and returns a [types.GithubConnector]
// ready for use. InitGithubConnector must be used to initialize any
// uninitialized [types.GithubConnector]s before they can be used.
func InitGithubConnector(c types.GithubConnector) (types.GithubConnector, error) {
githubConnectorMutex.RLock()
defer githubConnectorMutex.RUnlock()
return githubAuthInitializer(c)
}
// GithubAuthConverter converts a GitHub auth connector so it can be
// sent over gRPC.
type GithubAuthConverter func(types.GithubConnector) (*types.GithubConnectorV3, error)
var githubAuthConverter GithubAuthConverter
// RegisterGithubAuthConverter registers a function to convert GitHub auth connectors.
func RegisterGithubAuthConverter(convert GithubAuthConverter) {
githubConnectorMutex.Lock()
defer githubConnectorMutex.Unlock()
githubAuthConverter = convert
}
// ConvertGithubConnector converts a GitHub auth connector so it can be
// sent over gRPC.
func ConvertGithubConnector(c types.GithubConnector) (*types.GithubConnectorV3, error) {
githubConnectorMutex.RLock()
defer githubConnectorMutex.RUnlock()
return githubAuthConverter(c)
}
func init() {
RegisterGithubAuthCreator(types.NewGithubConnector)
RegisterGithubAuthInitializer(func(c types.GithubConnector) (types.GithubConnector, error) {
return c, nil
})
RegisterGithubAuthConverter(func(c types.GithubConnector) (*types.GithubConnectorV3, error) {
connector, ok := c.(*types.GithubConnectorV3)
if !ok {
return nil, trace.BadParameter("unrecognized github connector version %T", c)
}
return connector, nil
})
}
// UnmarshalGithubConnector unmarshals the GithubConnector resource from JSON.
func UnmarshalGithubConnector(bytes []byte, opts ...MarshalOption) (types.GithubConnector, error) {
r, err := UnmarshalResource(types.KindGithubConnector, bytes, opts...)
if err != nil {
return nil, err
}
connector, ok := r.(types.GithubConnector)
if !ok {
return nil, trace.BadParameter("expected GithubConnector, got %T", r)
}
return connector, nil
}
// UnmarshalOSSGithubConnector unmarshals the open source variant of the GithubConnector resource from JSON.
func UnmarshalOSSGithubConnector(bytes []byte, opts ...MarshalOption) (types.GithubConnector, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := json.Unmarshal(bytes, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var c types.GithubConnectorV3
if err := utils.FastUnmarshal(bytes, &c); err != nil {
return nil, trace.Wrap(err)
}
if err := c.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
c.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
c.SetExpiry(cfg.Expires)
}
return &c, nil
}
return nil, trace.BadParameter(
"GitHub connector resource version %q is not supported", h.Version)
}
// MarshalGithubConnector marshals a GithubConnector resource to JSON.
func MarshalGithubConnector(connector types.GithubConnector, opts ...MarshalOption) ([]byte, error) {
return MarshalResource(connector, opts...)
}
// MarshalOSSGithubConnector marshals the open source variant of the GithubConnector resource to JSON.
func MarshalOSSGithubConnector(githubConnector types.GithubConnector, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch githubConnector := githubConnector.(type) {
case *types.GithubConnectorV3:
if err := githubConnector.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
// Return an error for OSS build if the endpoint url is set, but it is
// not the public GitHub endpoint. Empty endpoint url is also allowed.
//
// Note that the enterprise marshaler also calls this marshaler to
// produce the final output.
if modules.GetModules().IsOSSBuild() {
if githubConnector.Spec.EndpointURL != "" &&
githubConnector.Spec.EndpointURL != types.GithubURL {
return nil, fmt.Errorf("GitHub endpoint URL is set: %w", ErrRequiresEnterprise)
}
if githubConnector.Spec.APIEndpointURL != "" &&
githubConnector.Spec.APIEndpointURL != types.GithubAPIURL {
return nil, fmt.Errorf("GitHub API endpoint URL is set: %w", ErrRequiresEnterprise)
}
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, githubConnector))
default:
return nil, trace.BadParameter("unrecognized github connector version %T", githubConnector)
}
}
/*
* 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 (
"crypto/sha256"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
)
// HeadlessAuthenticationUserStubID is the ID of a headless authentication stub.
const HeadlessAuthenticationUserStubID = "stub"
// ValidateHeadlessAuthentication verifies that the headless authentication has
// all of the required fields set. Headless authentication stubs will not pass
// this validation.
func ValidateHeadlessAuthentication(h *types.HeadlessAuthentication) error {
if err := h.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
switch {
case h.State.IsUnspecified():
return trace.BadParameter("headless authentication resource state must be specified")
case h.Version != types.V1:
return trace.BadParameter("unsupported headless authentication resource version %q, current supported version is %s", h.Version, types.V1)
case len(h.SshPublicKey) == 0:
return trace.BadParameter("headless authentication resource must have non-empty SSH public key")
case h.Metadata.Name != NewHeadlessAuthenticationID(h.SshPublicKey):
return trace.BadParameter("headless authentication resource name must be derived from public key")
}
return nil
}
// NewHeadlessAuthenticationID returns a new SHA256 (Version 5) UUID
// based on the supplied ssh public key.
func NewHeadlessAuthenticationID(pubKey []byte) string {
return uuid.NewHash(sha256.New(), uuid.Nil, pubKey, 5).String()
}
/*
* Teleport
* Copyright (C) 2025 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/trace"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/defaults"
healthcheckconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/healthcheckconfig/v1"
labelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/label/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/services/label"
)
// HealthCheckConfigReader defines methods for reading health check config
// resources.
type HealthCheckConfigReader interface {
// GetHealthCheckConfig fetches a health check config by name.
GetHealthCheckConfig(ctx context.Context, name string) (*healthcheckconfigv1.HealthCheckConfig, error)
// ListHealthCheckConfigs lists health check configs with pagination.
ListHealthCheckConfigs(ctx context.Context, limit int, startKey string) ([]*healthcheckconfigv1.HealthCheckConfig, string, error)
}
// HealthCheckConfig is a service that manages
// [healthcheckconfigv1.HealthCheckConfig] resources.
type HealthCheckConfig interface {
HealthCheckConfigReader
// CreateHealthCheckConfig creates a new health check config.
CreateHealthCheckConfig(ctx context.Context, in *healthcheckconfigv1.HealthCheckConfig) (*healthcheckconfigv1.HealthCheckConfig, error)
// UpdateHealthCheckConfig updates an existing health check config.
UpdateHealthCheckConfig(ctx context.Context, in *healthcheckconfigv1.HealthCheckConfig) (*healthcheckconfigv1.HealthCheckConfig, error)
// UpsertHealthCheckConfig creates or updates a health check config.
UpsertHealthCheckConfig(ctx context.Context, in *healthcheckconfigv1.HealthCheckConfig) (*healthcheckconfigv1.HealthCheckConfig, error)
// DeleteHealthCheckConfig deletes a health check config.
DeleteHealthCheckConfig(ctx context.Context, name string) error
}
// ValidateHealthCheckConfig validates the given health check config.
func ValidateHealthCheckConfig(s *healthcheckconfigv1.HealthCheckConfig) error {
switch {
case s == nil:
return trace.BadParameter("object must not be nil")
case s.GetVersion() != types.V1:
return trace.BadParameter("only version %q is supported, got %q", types.V1, s.GetVersion())
case s.GetKind() != types.KindHealthCheckConfig:
return trace.BadParameter("kind must be %q, got %q", types.KindHealthCheckConfig, s.GetKind())
case !s.HasMetadata():
return trace.BadParameter("metadata is missing")
case s.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name is missing")
case !s.HasSpec():
return trace.BadParameter("spec is missing")
case !s.GetSpec().HasMatch():
return trace.BadParameter("spec.match is missing")
}
for _, label := range s.GetSpec().GetMatch().GetDbLabels() {
if err := validateLabel(label); err != nil {
return trace.BadParameter("invalid spec.db_labels: %v", err)
}
}
if expr := s.GetSpec().GetMatch().GetDbLabelsExpression(); len(expr) > 0 {
if _, err := label.ParseExpression(expr); err != nil {
return trace.BadParameter("invalid spec.db_labels_expression: %v", err)
}
}
for _, label := range s.GetSpec().GetMatch().GetKubernetesLabels() {
if err := validateLabel(label); err != nil {
return trace.BadParameter("invalid spec.kubernetes_labels: %v", err)
}
}
if expr := s.GetSpec().GetMatch().GetKubernetesLabelsExpression(); len(expr) > 0 {
if _, err := label.ParseExpression(expr); err != nil {
return trace.BadParameter("invalid spec.kubernetes_labels_expression: %v", err)
}
}
timeout := s.GetSpec().GetTimeout().AsDuration()
switch {
case timeout == 0:
timeout = defaults.HealthCheckTimeout
case timeout < constants.MinHealthCheckTimeout:
return trace.BadParameter("spec.timeout must be at least %s", constants.MinHealthCheckTimeout)
}
interval := s.GetSpec().GetInterval().AsDuration()
switch {
case interval == 0:
interval = defaults.HealthCheckInterval
case interval < constants.MinHealthCheckInterval:
return trace.BadParameter("spec.interval must be at least %s", constants.MinHealthCheckInterval)
case interval > constants.MaxHealthCheckInterval:
return trace.BadParameter("spec.interval must not be greater than %s", constants.MaxHealthCheckInterval)
}
if timeout > interval {
if s.GetSpec().GetTimeout().AsDuration() == 0 {
return trace.BadParameter("spec.interval (%s) must not be less than the default timeout (%s)", interval, defaults.HealthCheckTimeout)
}
if s.GetSpec().GetInterval().AsDuration() == 0 {
return trace.BadParameter("spec.timeout (%s) must not be greater than the default interval (%s)", timeout, defaults.HealthCheckInterval)
}
return trace.BadParameter("spec.timeout (%s) must not be greater than spec.interval (%s)", timeout, interval)
}
if s.GetSpec().GetHealthyThreshold() > constants.MaxHealthCheckHealthyThreshold {
return trace.BadParameter(
"spec.healthy_threshold (%v) must not be greater than %v",
s.GetSpec().GetHealthyThreshold(),
constants.MaxHealthCheckHealthyThreshold,
)
}
if s.GetSpec().GetUnhealthyThreshold() > constants.MaxHealthCheckUnhealthyThreshold {
return trace.BadParameter(
"spec.unhealthy_threshold (%v) must not be greater than %v",
s.GetSpec().GetUnhealthyThreshold(),
constants.MaxHealthCheckUnhealthyThreshold,
)
}
return nil
}
func validateLabel(label *labelv1.Label) error {
if label.GetName() == types.Wildcard {
if len(label.GetValues()) != 1 || label.GetValues()[0] != types.Wildcard {
return trace.BadParameter("selector *:%s is not supported, a wildcard label key may only be used with a wildcard label value", label.GetValues()[0])
}
}
return nil
}
// MarshalHealthCheckConfig marshals HealthCheckConfig resource to JSON.
func MarshalHealthCheckConfig(cfg *healthcheckconfigv1.HealthCheckConfig, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(cfg, opts...)
}
// UnmarshalHealthCheckConfig unmarshals the HealthCheckConfig resource.
func UnmarshalHealthCheckConfig(data []byte, opts ...MarshalOption) (*healthcheckconfigv1.HealthCheckConfig, error) {
return UnmarshalProtoResource[*healthcheckconfigv1.HealthCheckConfig](data, opts...)
}
/*
* 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 implements API services exposed by Teleport:
// * presence service that takes care of heartbeats
// * web service that takes care of web logins
// * ca service - certificate authorities
package services
import (
"context"
"crypto"
"crypto/x509"
"iter"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/proto"
userspb "github.com/gravitational/teleport/api/gen/proto/go/teleport/users/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/keys"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/defaults"
)
// UserGetter is responsible for getting users
type UserGetter interface {
// GetUser returns a user by name
GetUser(ctx context.Context, user string, withSecrets bool) (types.User, error)
}
// UsersService is responsible for basic user management
type UsersService interface {
UserGetter
// UpdateUser updates an existing user.
UpdateUser(ctx context.Context, user types.User) (types.User, error)
// UpdateAndSwapUser reads an existing user, runs `fn` against it and writes
// the result to storage. Return `false` from `fn` to avoid storage changes.
// Roughly equivalent to [GetUser] followed by [CompareAndSwapUser].
// Returns the storage user.
UpdateAndSwapUser(ctx context.Context, user string, withSecrets bool, fn func(types.User) (changed bool, err error)) (types.User, error)
// UpsertUser updates parameters about user
UpsertUser(ctx context.Context, user types.User) (types.User, error)
// CompareAndSwapUser updates an existing user, but fails if the user does
// not match an expected backend value.
CompareAndSwapUser(ctx context.Context, new, existing types.User) error
// DeleteUser deletes a user with all the keys from the backend
DeleteUser(ctx context.Context, user string) error
// GetUsers returns a list of users registered with the local auth server
GetUsers(ctx context.Context, withSecrets bool) ([]types.User, error)
// ListUsers returns a page of users.
ListUsers(ctx context.Context, req *userspb.ListUsersRequest) (*userspb.ListUsersResponse, error)
}
// IdentityInternal extends the Identity interface with auth-specific internal methods.
type IdentityInternal interface {
Identity
// AppendPutUserParamsActions adds conditional actions to an atomic write to
// create or update the user params resource (without secrets, mfa devices).
AppendPutUserParamsActions(
actions []backend.ConditionalAction,
user types.User,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteUserParamsActions adds conditional actions to an atomic write
// to delete the user params resource.
//
// Note: the returned actions will NOT delete the user's password, MFA devices,
// etc. so is only really suitable for bot users, in most cases you should use
// DeleteUser instead.
AppendDeleteUserParamsActions(
actions []backend.ConditionalAction,
user string,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
// Identity is responsible for managing user entries and external identities
type Identity interface {
// CreateUser creates user, only if the user entry does not exist
CreateUser(ctx context.Context, user types.User) (types.User, error)
// UsersService implements most methods
UsersService
// AddUserLoginAttempt logs user login attempt
AddUserLoginAttempt(user string, attempt LoginAttempt, ttl time.Duration) error
// GetUserLoginAttempts returns user login attempts
GetUserLoginAttempts(user string) ([]LoginAttempt, error)
// DeleteUserLoginAttempts removes all login attempts of a user. Should be
// called after successful login.
DeleteUserLoginAttempts(user string) error
// GetUserByOIDCIdentity returns a user by its specified OIDC Identity, returns first
// user specified with this identity
GetUserByOIDCIdentity(id types.ExternalIdentity) (types.User, error)
// GetUserBySAMLIdentity returns a user by its specified OIDC Identity, returns first
// user specified with this identity
GetUserBySAMLIdentity(id types.ExternalIdentity) (types.User, error)
// GetUserByGithubIdentity returns a user by its specified Github identity
GetUserByGithubIdentity(id types.ExternalIdentity) (types.User, error)
// GetPasswordHash returns the password hash for a given user
GetPasswordHash(user string) ([]byte, error)
// UpsertUsedTOTPToken upserts a TOTP token to the backend so it can't be used again
// during the 30 second window it's valid.
UpsertUsedTOTPToken(user string, otpToken string) error
// GetUsedTOTPToken returns the last successfully used TOTP token.
GetUsedTOTPToken(user string) (string, error)
// UpsertPassword upserts a new password. It also sets the user's
// `PasswordState` status flag accordingly. Returns an error if the user
// doesn't exist.
UpsertPassword(user string, password []byte) error
// DeletePassword deletes user's password and sets the `PasswordState` status
// flag accordingly.
DeletePassword(ctx context.Context, username string) error
// UpsertWebauthnLocalAuth creates or updates the local auth configuration for
// Webauthn.
// WebauthnLocalAuth is a component of LocalAuthSecrets.
// Automatically indexes the WebAuthn user ID for lookup by
// GetTeleportUserByWebauthnID.
UpsertWebauthnLocalAuth(ctx context.Context, user string, wla *types.WebauthnLocalAuth) error
// GetWebauthnLocalAuth retrieves the existing local auth configuration for
// Webauthn, if any.
// WebauthnLocalAuth is a component of LocalAuthSecrets.
GetWebauthnLocalAuth(ctx context.Context, user string) (*types.WebauthnLocalAuth, error)
// GetTeleportUserByWebauthnID reads a Teleport username from a WebAuthn user
// ID (aka user handle).
// See UpsertWebauthnLocalAuth and types.WebauthnLocalAuth.
GetTeleportUserByWebauthnID(ctx context.Context, webID []byte) (string, error)
// UpsertWebauthnSessionData creates or updates WebAuthn session data in
// storage, for the purpose of later verifying an authentication or
// registration challenge.
// Session data is expected to expire according to backend settings.
UpsertWebauthnSessionData(ctx context.Context, user, sessionID string, sd *wantypes.SessionData) error
// GetWebauthnSessionData retrieves a previously-stored session data by ID,
// if it exists and has not expired.
GetWebauthnSessionData(ctx context.Context, user, sessionID string) (*wantypes.SessionData, error)
// DeleteWebauthnSessionData deletes session data by ID, if it exists and has
// not expired.
DeleteWebauthnSessionData(ctx context.Context, user, sessionID string) error
// UpsertGlobalWebauthnSessionData creates or updates WebAuthn session data in
// storage, for the purpose of later verifying an authentication challenge.
// Session data is expected to expire according to backend settings.
// Used for passwordless challenges.
UpsertGlobalWebauthnSessionData(ctx context.Context, scope, id string, sd *wantypes.SessionData) error
// GetGlobalWebauthnSessionData retrieves previously-stored session data by ID,
// if it exists and has not expired.
// Used for passwordless challenges.
GetGlobalWebauthnSessionData(ctx context.Context, scope, id string) (*wantypes.SessionData, error)
// DeleteGlobalWebauthnSessionData deletes session data by ID, if it exists
// and has not expired.
DeleteGlobalWebauthnSessionData(ctx context.Context, scope, id string) error
// UpsertMFADevice upserts an MFA device for the user.
UpsertMFADevice(ctx context.Context, user string, d *types.MFADevice) error
// GetMFADevices gets all MFA devices for the user.
GetMFADevices(ctx context.Context, user string, withSecrets bool) ([]*types.MFADevice, error)
// DeleteMFADevice deletes an MFA device for the user by ID.
DeleteMFADevice(ctx context.Context, user, id string) error
// CreateOIDCConnector creates a new OIDC connector.
CreateOIDCConnector(ctx context.Context, connector types.OIDCConnector) (types.OIDCConnector, error)
// UpdateOIDCConnector updates an existing OIDC connector.
UpdateOIDCConnector(ctx context.Context, connector types.OIDCConnector) (types.OIDCConnector, error)
// UpsertOIDCConnector updates or creates an OIDC connector.
UpsertOIDCConnector(ctx context.Context, connector types.OIDCConnector) (types.OIDCConnector, error)
// DeleteOIDCConnector deletes OIDC Connector
DeleteOIDCConnector(ctx context.Context, connectorID string) error
// GetOIDCConnector returns OIDC connector data, withSecrets adds or removes client secret from return results
GetOIDCConnector(ctx context.Context, id string, withSecrets bool) (types.OIDCConnector, error)
// GetOIDCConnectors returns valid registered connectors, withSecrets adds or removes client secret from return
// results. Invalid Connectors are simply logged but errors are not forwarded.
GetOIDCConnectors(ctx context.Context, withSecrets bool) ([]types.OIDCConnector, error)
// ListOIDCConnectors returns a page of valid registered connectors.
// withSecrets adds or removes client secret from return results.
ListOIDCConnectors(ctx context.Context, limit int, start string, withSecrets bool) ([]types.OIDCConnector, string, error)
// RangeOIDCConnectors returns valid registered connectors within the range [start, end).
// withSecrets adds or removes client secret from return results.
RangeOIDCConnectors(ctx context.Context, start, end string, withSecrets bool) iter.Seq2[types.OIDCConnector, error]
// CreateOIDCAuthRequest creates new auth request
CreateOIDCAuthRequest(ctx context.Context, req types.OIDCAuthRequest, ttl time.Duration) error
// GetOIDCAuthRequest returns OIDC auth request if found
GetOIDCAuthRequest(ctx context.Context, stateToken string) (*types.OIDCAuthRequest, error)
// CreateSAMLConnector creates a new SAML connector.
CreateSAMLConnector(ctx context.Context, connector types.SAMLConnector) (types.SAMLConnector, error)
// UpdateSAMLConnector updates an existing SAML connector
UpdateSAMLConnector(ctx context.Context, connector types.SAMLConnector) (types.SAMLConnector, error)
// UpsertSAMLConnector updates or creates a SAML connector
UpsertSAMLConnector(ctx context.Context, connector types.SAMLConnector) (types.SAMLConnector, error)
// DeleteSAMLConnector deletes OIDC Connector
DeleteSAMLConnector(ctx context.Context, connectorID string) error
// GetSAMLConnector returns OIDC connector data, withSecrets adds or removes secrets from return results
GetSAMLConnector(ctx context.Context, id string, withSecrets bool) (types.SAMLConnector, error)
// GetSAMLConnector returns OIDC connector data, withSecrets adds or removes secrets from return results
GetSAMLConnectorWithValidationOptions(ctx context.Context, id string, withSecrets bool, opts ...types.SAMLConnectorValidationOption) (types.SAMLConnector, error)
// GetSAMLConnectors returns valid registered connectors, withSecrets adds or removes secret from return results.
// Invalid Connectors are simply logged but errors are not forwarded.
GetSAMLConnectors(ctx context.Context, withSecrets bool) ([]types.SAMLConnector, error)
// GetSAMLConnectors returns valid registered connectors, withSecrets adds or removes secret from return results.
// Invalid Connectors are simply logged but errors are not forwarded.
GetSAMLConnectorsWithValidationOptions(ctx context.Context, withSecrets bool, opts ...types.SAMLConnectorValidationOption) ([]types.SAMLConnector, error)
// ListSAMLConnectorsWithOptions returns a page of valid registered connectors.
// withSecrets adds or removes client secret from return results.
ListSAMLConnectorsWithOptions(ctx context.Context, limit int, start string, withSecrets bool, opts ...types.SAMLConnectorValidationOption) ([]types.SAMLConnector, string, error)
// RangeSAMLConnectorsWithOptions returns valid registered connectors within the range [start, end).
// withSecrets adds or removes client secret from return results.
RangeSAMLConnectorsWithOptions(ctx context.Context, start, end string, withSecrets bool, opts ...types.SAMLConnectorValidationOption) iter.Seq2[types.SAMLConnector, error]
// CreateSAMLAuthRequest creates new auth request
CreateSAMLAuthRequest(ctx context.Context, req types.SAMLAuthRequest, ttl time.Duration) error
// GetSAMLAuthRequest returns SAML auth request if found
GetSAMLAuthRequest(ctx context.Context, id string) (*types.SAMLAuthRequest, error)
// CreateSSODiagnosticInfo creates new SSO diagnostic info record.
CreateSSODiagnosticInfo(ctx context.Context, authKind string, authRequestID string, entry types.SSODiagnosticInfo) error
// GetSSODiagnosticInfo returns SSO diagnostic info records.
GetSSODiagnosticInfo(ctx context.Context, authKind string, authRequestID string) (*types.SSODiagnosticInfo, error)
// CreateGithubConnector creates a new Github connector.
CreateGithubConnector(ctx context.Context, connector types.GithubConnector) (types.GithubConnector, error)
// UpdateGithubConnector updates an existing Github connector.
UpdateGithubConnector(ctx context.Context, connector types.GithubConnector) (types.GithubConnector, error)
// UpsertGithubConnector creates or updates a Github connector.
UpsertGithubConnector(ctx context.Context, connector types.GithubConnector) (types.GithubConnector, error)
// GetGithubConnectors returns valid Github connectors, invalid Connectors are simply logged but errors are not forwarded.
GetGithubConnectors(ctx context.Context, withSecrets bool) ([]types.GithubConnector, error)
// ListGithubConnectors returns a page of valid registered Github connectors.
// withSecrets adds or removes client secret from return results.
ListGithubConnectors(ctx context.Context, limit int, start string, withSecrets bool) ([]types.GithubConnector, string, error)
// RangeGithubConnectors returns valid registered Github connectors within the range [start, end).
// withSecrets adds or removes client secret from return results.
RangeGithubConnectors(ctx context.Context, start, end string, withSecrets bool) iter.Seq2[types.GithubConnector, error]
// GetGithubConnector returns a Github connector by its name
GetGithubConnector(ctx context.Context, name string, withSecrets bool) (types.GithubConnector, error)
// DeleteGithubConnector deletes a Github connector by its name
DeleteGithubConnector(ctx context.Context, name string) error
// CreateGithubAuthRequest creates a new auth request for Github OAuth2 flow
CreateGithubAuthRequest(ctx context.Context, req types.GithubAuthRequest) error
// GetGithubAuthRequest retrieves Github auth request by the token
GetGithubAuthRequest(ctx context.Context, stateToken string) (*types.GithubAuthRequest, error)
// UpsertMFASessionData creates or updates MFA session data in
// storage, for the purpose of later verifying an MFA authentication attempt.
// MFA session data is expected to expire according to backend settings.
// Used for both SSO and Browser MFA.
UpsertMFASessionData(ctx context.Context, sd *MFASessionData) error
// GetMFASessionData retrieves SSO or Browser MFA session data by ID.
GetMFASessionData(ctx context.Context, sessionID string) (*MFASessionData, error)
// DeleteMFASessionData deletes SSO or Browser MFA session data by ID.
DeleteMFASessionData(ctx context.Context, sessionID string) error
// TODO(danielashare): Remove deprecated *SSOMFASessionData functions once teleport.e is using the new functions
// UpsertSSOMFASessionData creates or updates SSO MFA session data in
// storage, for the purpose of later verifying an SSO MFA authentication
// attempt.
//
// Deprecated: use UpsertMFASessionData.
UpsertSSOMFASessionData(ctx context.Context, sd *SSOMFASessionData) error
// GetSSOMFASessionData retrieves SSO MFA session data by ID.
//
// Deprecated: use GetMFASessionData.
GetSSOMFASessionData(ctx context.Context, sessionID string) (*SSOMFASessionData, error)
// DeleteSSOMFASessionData deletes SSO MFA session data by ID.
//
// Deprecated: use DeleteMFASessionData.
DeleteSSOMFASessionData(ctx context.Context, sessionID string) error
// CreateUserToken creates a new user token.
CreateUserToken(ctx context.Context, token types.UserToken) (types.UserToken, error)
// DeleteUserToken deletes a user token.
DeleteUserToken(ctx context.Context, tokenID string) error
// ListUserTokens returns a page of user tokens.
ListUserTokens(ctx context.Context, limit int, startKey string) ([]types.UserToken, string, error)
// GetUserToken returns a user token by id.
GetUserToken(ctx context.Context, tokenID string) (types.UserToken, error)
// UpsertUserTokenSecrets upserts a user token secrets.
UpsertUserTokenSecrets(ctx context.Context, secrets types.UserTokenSecrets) error
// GetUserTokenSecrets returns a user token secrets.
GetUserTokenSecrets(ctx context.Context, tokenID string) (types.UserTokenSecrets, error)
// UpsertRecoveryCodes upserts a user's new recovery codes.
UpsertRecoveryCodes(ctx context.Context, user string, recovery *types.RecoveryCodesV1) error
// GetRecoveryCodes gets a user's recovery codes.
GetRecoveryCodes(ctx context.Context, user string, withSecrets bool) (*types.RecoveryCodesV1, error)
// UpsertKeyAttestationData upserts a verified public key attestation response.
UpsertKeyAttestationData(ctx context.Context, attestationData *keys.AttestationData, ttl time.Duration) error
// GetKeyAttestationData gets a verified public key attestation response.
GetKeyAttestationData(ctx context.Context, pubDer []byte) (*keys.AttestationData, error)
HeadlessAuthenticationService
types.WebSessionsGetter
WebToken
// AppSession defines application session features.
AppSession
// SnowflakeSession defines Snowflake session features.
SnowflakeSession
}
// AppSessionReader defines application session features available to remote clients.
type AppSessionReader interface {
// GetAppSession gets an application web session.
GetAppSession(context.Context, types.GetAppSessionRequest) (types.WebSession, error)
// ListAppSessions gets a paginated list of application web sessions.
ListAppSessions(ctx context.Context, pageSize int, pageToken, user string) ([]types.WebSession, string, error)
// DeleteAppSession removes an application web session.
DeleteAppSession(context.Context, types.DeleteAppSessionRequest) error
// DeleteAllAppSessions removes all application web sessions.
DeleteAllAppSessions(context.Context) error
// DeleteUserAppSessions deletes all user’s application sessions.
DeleteUserAppSessions(ctx context.Context, req *proto.DeleteUserAppSessionsRequest) error
}
// AppSession defines application session features.
type AppSession interface {
AppSessionReader
// UpdateAppSession updates an existing application web session if the revisions match.
UpdateAppSession(context.Context, types.WebSession) error
// UpsertAppSession upserts an application web session.
UpsertAppSession(context.Context, types.WebSession) error
}
// SnowflakeSession defines Snowflake session features.
type SnowflakeSession interface {
// GetSnowflakeSession gets a Snowflake web session.
GetSnowflakeSession(context.Context, types.GetSnowflakeSessionRequest) (types.WebSession, error)
// GetSnowflakeSessions gets all Snowflake web sessions.
GetSnowflakeSessions(context.Context) ([]types.WebSession, error)
// ListSnowflakeSessions returns a page of Snowflake web sessions.
ListSnowflakeSessions(ctx context.Context, limit int, start string) ([]types.WebSession, string, error)
// RangeSnowflakeSessions returns Snowflake web sessions within the range [start, end).
RangeSnowflakeSessions(ctx context.Context, start, end string) iter.Seq2[types.WebSession, error]
// UpsertSnowflakeSession upserts a Snowflake web session.
UpsertSnowflakeSession(context.Context, types.WebSession) error
// DeleteSnowflakeSession removes a Snowflake web session.
DeleteSnowflakeSession(context.Context, types.DeleteSnowflakeSessionRequest) error
// DeleteAllSnowflakeSessions removes all Snowflake web sessions.
DeleteAllSnowflakeSessions(context.Context) error
}
// HeadlessAuthenticationService is responsible for headless authentication resource management
type HeadlessAuthenticationService interface {
// GetHeadlessAuthentication gets a headless authentication.
GetHeadlessAuthentication(ctx context.Context, username, name string) (*types.HeadlessAuthentication, error)
// GetHeadlessAuthentications gets all headless authentications.
GetHeadlessAuthentications(ctx context.Context) ([]*types.HeadlessAuthentication, error)
// UpsertHeadlessAuthentication upserts a headless authentication.
UpsertHeadlessAuthentication(ctx context.Context, ha *types.HeadlessAuthentication) error
// CompareAndSwapHeadlessAuthentication performs a compare
// and swap replacement on a headless authentication resource.
CompareAndSwapHeadlessAuthentication(ctx context.Context, old, new *types.HeadlessAuthentication) (*types.HeadlessAuthentication, error)
// DeleteHeadlessAuthentication deletes a headless authentication from the backend.
DeleteHeadlessAuthentication(ctx context.Context, username, name string) error
// DeleteAllHeadlessAuthentications deletes all headless authentications from the backend.
DeleteAllHeadlessAuthentications(ctx context.Context) error
}
// VerifyPassword makes sure password satisfies our requirements (relaxed),
// mostly to avoid putting garbage in
func VerifyPassword(password []byte) error {
if len(password) < defaults.MinPasswordLength {
return trace.BadParameter(
"password is too short, min length is %v", defaults.MinPasswordLength)
}
if len(password) > defaults.MaxPasswordLength {
return trace.BadParameter(
"password is too long, max length is %v", defaults.MaxPasswordLength)
}
return nil
}
// Users represents a slice of users,
// makes it sort compatible (sorts by username)
type Users []types.User
func (u Users) Len() int {
return len(u)
}
func (u Users) Less(i, j int) bool {
return u[i].GetName() < u[j].GetName()
}
func (u Users) Swap(i, j int) {
u[i], u[j] = u[j], u[i]
}
// SortedLoginAttempts sorts login attempts by time
type SortedLoginAttempts []LoginAttempt
// Len returns length of a role list
func (s SortedLoginAttempts) Len() int {
return len(s)
}
// Less stacks latest attempts to the end of the list
func (s SortedLoginAttempts) Less(i, j int) bool {
return s[i].Time.Before(s[j].Time)
}
// Swap swaps two attempts
func (s SortedLoginAttempts) Swap(i, j int) {
s[i], s[j] = s[j], s[i]
}
// LastFailed calculates last x successive attempts are failed
func LastFailed(x int, attempts []LoginAttempt) bool {
var failed int
for i := len(attempts) - 1; i >= 0; i-- {
if !attempts[i].Success {
failed++
} else {
return false
}
if failed >= x {
return true
}
}
return false
}
// NewWebSessionAttestationData creates attestation data for a web session key.
// Inserting data to the Auth server will allow certificates generated for the
// web session key to pass private key policies that are unobtainable in the web
// (hardware key policies). In exchange, these keys must be kept strictly in the
// Auth and Proxy processes and Auth storage. These keys and certs can only be
// retrieved by users in the form of web session cookies.
func NewWebSessionAttestationData(pub crypto.PublicKey) (*keys.AttestationData, error) {
pubDER, err := x509.MarshalPKIXPublicKey(pub)
if err != nil {
return nil, trace.Wrap(err)
}
return &keys.AttestationData{
PublicKeyDER: pubDER,
PrivateKeyPolicy: keys.PrivateKeyPolicyWebSession,
}, nil
}
// Teleport
// Copyright (C) 2024 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"
"fmt"
"github.com/gravitational/trace"
identitycenterv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/identitycenter/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// IdentityCenterAccount wraps a raw identity center record in a new type to
// allow it to implement the interfaces required for use with the Unified
// Resource listing.
//
// IdentityCenterAccount simply wraps a pointer to the underlying
// identitycenterv1.Account record, and can be treated as a reference-like type.
// Copies of an IdentityCenterAccount will point to the same record.
type IdentityCenterAccount struct {
// This wrapper needs to:
// - implement the interfaces required for use with the Unified Resource
// service.
// - expose the existing interfaces & methods on the underlying
// identitycenterv1.Account
// - avoid copying the underlying identitycenterv1.Account due to embedded
// mutexes in the protobuf-generated code
//
// Given those requirements, storing an embedded pointer seems to be the
// least-bad approach.
*identitycenterv1.Account
}
// GetDisplayName returns a human-readable name for the account for UI display.
func (a IdentityCenterAccount) GetDisplayName() string {
return a.Account.GetSpec().GetName()
}
// IdentityCenterAccountID is a strongly-typed Identity Center account ID.
type IdentityCenterAccountID string
// IdentityCenterAccountGetter provides read-only access to Identity Center
// Account records
type IdentityCenterAccountGetter interface {
// ListIdentityCenterAccounts provides a paged list of all known identity
// center accounts
ListIdentityCenterAccounts(context.Context, int, string) ([]*identitycenterv1.Account, string, error)
// GetIdentityCenterAccount fetches a specific Identity Center Account
GetIdentityCenterAccount(context.Context, string) (*identitycenterv1.Account, error)
}
// IdentityCenterAccounts defines read/write access to Identity Center account
// resources
type IdentityCenterAccounts interface {
IdentityCenterAccountGetter
// CreateIdentityCenterAccount creates a new Identity Center Account record
CreateIdentityCenterAccount(context.Context, *identitycenterv1.Account) (*identitycenterv1.Account, error)
// UpdateIdentityCenterAccount performs a conditional update on an Identity
// Center Account record, returning the updated record on success.
UpdateIdentityCenterAccount(context.Context, *identitycenterv1.Account) (*identitycenterv1.Account, error)
// UpsertIdentityCenterAccount performs an *unconditional* upsert on an
// Identity Center Account record, returning the updated record on success.
// Be careful when mixing UpsertIdentityCenterAccount() with resources
// protected by optimistic locking
UpsertIdentityCenterAccount(context.Context, *identitycenterv1.Account) (*identitycenterv1.Account, error)
// DeleteIdentityCenterAccount deletes an Identity Center Account record
DeleteIdentityCenterAccount(context.Context, IdentityCenterAccountID) error
// DeleteAllIdentityCenterAccounts deletes all Identity Center Account records
DeleteAllIdentityCenterAccounts(context.Context) error
}
// PrincipalAssignmentID is a strongly-typed ID for Identity Center Principal
// Assignments
type PrincipalAssignmentID string
// IdentityCenterPrincipalAssignments defines operations on an Identity Center
// principal assignment database
type IdentityCenterPrincipalAssignments interface {
// ListPrincipalAssignments lists all PrincipalAssignment records in the
// service
ListPrincipalAssignments(context.Context, int, string) ([]*identitycenterv1.PrincipalAssignment, string, error)
// CreatePrincipalAssignment creates a new Principal Assignment record in
// the service from the supplied in-memory representation. Returns the
// created record on success.
CreatePrincipalAssignment(context.Context, *identitycenterv1.PrincipalAssignment) (*identitycenterv1.PrincipalAssignment, error)
// GetPrincipalAssignment fetches a specific Principal Assignment record.
GetPrincipalAssignment(context.Context, PrincipalAssignmentID) (*identitycenterv1.PrincipalAssignment, error)
// UpdatePrincipalAssignment performs a conditional update on a Principal
// Assignment record
UpdatePrincipalAssignment(context.Context, *identitycenterv1.PrincipalAssignment) (*identitycenterv1.PrincipalAssignment, error)
// UpsertPrincipalAssignment performs an unconditional update on a Principal
// Assignment record
UpsertPrincipalAssignment(context.Context, *identitycenterv1.PrincipalAssignment) (*identitycenterv1.PrincipalAssignment, error)
// DeletePrincipalAssignment deletes a specific principal assignment record
DeletePrincipalAssignment(context.Context, PrincipalAssignmentID) error
// DeleteAllPrincipalAssignments deletes all assignment record
DeleteAllPrincipalAssignments(context.Context) error
}
// PermissionSetID is a strongly typed ID for an identitycenterv1.PermissionSet
type PermissionSetID string
// IdentityCenterPermissionSets defines the operations to create and maintain
// identitycenterv1.PermissionSet records in the service.
type IdentityCenterPermissionSets interface {
// ListPermissionSets list the known Permission Sets
ListPermissionSets(context.Context, int, string) ([]*identitycenterv1.PermissionSet, string, error)
// CreatePermissionSet creates a new PermissionSet record based on the
// supplied in-memory representation, returning the created record on
// success
CreatePermissionSet(context.Context, *identitycenterv1.PermissionSet) (*identitycenterv1.PermissionSet, error)
// GetPermissionSet fetches a specific PermissionSet record
GetPermissionSet(context.Context, PermissionSetID) (*identitycenterv1.PermissionSet, error)
// UpdatePermissionSet performs a conditional update on the supplied Identity
// Center Permission Set
UpdatePermissionSet(context.Context, *identitycenterv1.PermissionSet) (*identitycenterv1.PermissionSet, error)
// DeletePermissionSet deletes a specific Identity Center PermissionSet
DeletePermissionSet(context.Context, PermissionSetID) error
// DeleteAllPermissionSets deletes all Identity Center PermissionSets.
DeleteAllPermissionSets(context.Context) error
}
// IdentityCenterAccountAssignment wraps a raw identitycenterv1.AccountAssignment
// record in a new type to allow it to implement the interfaces required for use
// with the Unified Resource listing. IdentityCenterAccountAssignment simply
// wraps a pointer to the underlying account record, and can be treated as a
// reference-like type.
//
// Copies of an IdentityCenterAccountAssignment will point to the same record.
type IdentityCenterAccountAssignment struct {
// This wrapper needs to:
// - implement the interfaces required for use with the Unified Resource
// service.
// - expose the existing interfaces & methods on the underlying
// identitycenterv1.AccountAssignment
// - avoid copying the underlying identitycenterv1.AccountAssignment due to
// embedded mutexes in the protobuf-generated code
//
// Given those requirements, storing an embedded pointer seems to be the
// least-bad approach.
*identitycenterv1.AccountAssignment
}
// IdentityCenterAccountAssignmentID is a strongly typed ID for an
// IdentityCenterAccountAssignment
type IdentityCenterAccountAssignmentID string
// IdentityCenterAccountAssignmentGetter provides read-only access to Identity
// Center Account Assignment records
type IdentityCenterAccountAssignmentGetter interface {
// GetIdentityCenterAccountAssignment fetches a specific Account Assignment record.
GetIdentityCenterAccountAssignment(context.Context, string) (*identitycenterv1.AccountAssignment, error)
// ListIdentityCenterAccountAssignments provides a page of AccountAssignment records.
ListIdentityCenterAccountAssignments(context.Context, int, string) ([]*identitycenterv1.AccountAssignment, string, error)
}
// IdentityCenterAccountAssignments defines the operations to create and maintain
// Identity Center account assignment records in the service.
type IdentityCenterAccountAssignments interface {
IdentityCenterAccountAssignmentGetter
// CreateIdentityCenterAccountAssignment creates a new Account Assignment record in
// the service from the supplied in-memory representation. Returns the
// created record on success.
CreateIdentityCenterAccountAssignment(context.Context, *identitycenterv1.AccountAssignment) (*identitycenterv1.AccountAssignment, error)
// UpdateIdentityCenterAccountAssignment performs a conditional update on the supplied
// Account Assignment, returning the updated record on success.
UpdateIdentityCenterAccountAssignment(context.Context, *identitycenterv1.AccountAssignment) (*identitycenterv1.AccountAssignment, error)
// UpsertIdentityCenterAccountAssignment performs an unconditional update on the supplied
// Account Assignment, returning the updated record on success.
UpsertIdentityCenterAccountAssignment(context.Context, *identitycenterv1.AccountAssignment) (*identitycenterv1.AccountAssignment, error)
// DeleteIdentityCenterAccountAssignment deletes a specific account assignment
DeleteIdentityCenterAccountAssignment(context.Context, IdentityCenterAccountAssignmentID) error
// DeleteAllIdentityCenterAccountAssignments deletes all known account assignments
DeleteAllIdentityCenterAccountAssignments(context.Context) error
// DeleteAccountAssignment deletes a specific account assignment
// Deprecated: Prefer using DeleteIdentityCenterAccountAssignment
DeleteAccountAssignment(context.Context, IdentityCenterAccountAssignmentID) error
// DeleteAllAccountAssignments deletes all known account assignments
// Deprecated: Prefer using DeleteAllIdentityCenterAccountAssignment
DeleteAllAccountAssignments(context.Context) error
}
// IdentityCenter combines all the resource managers used by the Identity Center plugin
type IdentityCenter interface {
IdentityCenterAccounts
IdentityCenterPermissionSets
IdentityCenterPrincipalAssignments
IdentityCenterAccountAssignments
}
func IdentityCenterAccountToAppServer(acct *identitycenterv1.Account) *types.AppServerV3 {
srcPSs := acct.GetSpec().GetPermissionSetInfo()
pss := make([]*types.IdentityCenterPermissionSet, len(srcPSs))
for i, ps := range acct.GetSpec().GetPermissionSetInfo() {
pss[i] = &types.IdentityCenterPermissionSet{
ARN: ps.GetArn(),
Name: ps.GetName(),
AssignmentID: ps.GetAssignmentId(),
}
}
// Identity Center accounts surface in the unified-resource cache as
// synthetic AppServers; they never traverse the app write paths
// (ValidateApp / ValidateAppServer), so no DNS-1123 normalization
// here. The web Launch button (ResourceActionButton.tsx) builds the
// SSO launch URL as `${publicAddr}&role_name=...`, which requires
// the full StartUrl - scheme, path, and case preserved.
metadata := types.Metadata153ToLegacy(acct.GetMetadata())
metadata.Description = acct.GetSpec().GetName()
return &types.AppServerV3{
Kind: types.KindAppServer,
SubKind: types.KindIdentityCenterAccount,
Version: types.V3,
Metadata: metadata,
Spec: types.AppServerSpecV3{
App: &types.AppV3{
Kind: types.KindApp,
SubKind: types.KindIdentityCenterAccount,
Version: types.V3,
Metadata: metadata,
Spec: types.AppSpecV3{
URI: acct.GetSpec().GetStartUrl(),
PublicAddr: acct.GetSpec().GetStartUrl(),
AWS: &types.AppAWS{
ExternalID: acct.GetSpec().GetId(),
},
IdentityCenter: &types.AppIdentityCenter{
AccountID: acct.GetSpec().GetId(),
PermissionSets: pss,
},
},
},
},
}
}
// NewIdentityCenterAppMatcher creates a new [RoleMatcher] configured to
// match the supplied [types.Application] that is wrapping a [*identitycenterv1.Account].
func NewIdentityCenterAppMatcher(app types.Application) *IdentityCenterAccountMatcher {
ic := app.GetIdentityCenter()
if ic == nil {
return nil
}
return &IdentityCenterAccountMatcher{accountID: ic.AccountID}
}
// NewIdentityCenterAccountMatcher creates a new [RoleMatcher] configured to
// match the supplied [IdentityCenterAccount].
func NewIdentityCenterAccountMatcher(account IdentityCenterAccount) *IdentityCenterAccountMatcher {
return &IdentityCenterAccountMatcher{
accountID: account.GetSpec().GetId(),
}
}
// IdentityCenterMatcher implements a [RoleMatcher] for comparing Identity Center
// Account resources against the AccountAssignments specified in a Role condition.
type IdentityCenterAccountMatcher struct {
accountID string
}
// Match implements Role Matching for Identity Center Account resources. It
// attempts to match the Account Assignments in a Role Condition against a
// known Account ID.
func (m *IdentityCenterAccountMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
// TODO(tcsc): Expand to cover role template expansion (e.g. {{external.account_assignments}})
for _, asmt := range role.GetIdentityCenterAccountAssignments(condition) {
accountMatches, err := matchExpression(m.accountID, asmt.Account)
if err != nil {
return false, trace.Wrap(err)
}
if accountMatches {
return true, nil
}
}
return false, nil
}
func (m *IdentityCenterAccountMatcher) String() string {
return fmt.Sprintf("IdentityCenterAccountMatcher(account=%v)", m.accountID)
}
// NewIdentityCenterAccountAssignmentMatcher creates a new [IdentityCenterAccountAssignmentMatcher]
// configured to match the supplied [IdentityCenterAccountAssignment].
func NewIdentityCenterAccountAssignmentMatcher(assignment IdentityCenterAccountAssignment) *IdentityCenterAccountAssignmentMatcher {
return &IdentityCenterAccountAssignmentMatcher{
accountID: assignment.GetSpec().GetAccountId(),
permissionSetARN: assignment.GetSpec().GetPermissionSet().GetArn(),
}
}
// IdentityCenterMatcher implements a [RoleMatcher] for comparing Identity Center
// Account Assignment resources against the AccountAssignments specified in a
// Role condition.
type IdentityCenterAccountAssignmentMatcher struct {
accountID string
permissionSetARN string
}
// Match implements Role Matching for Identity Center Account resources. It
// attempts to match the Account Assignments in a Role Condition against a
// known Account ID.
func (m *IdentityCenterAccountAssignmentMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
// TODO(tcsc): Expand to cover role template expansion (e.g. {{external.account_assignments}})
for _, asmt := range role.GetIdentityCenterAccountAssignments(condition) {
accountMatches, err := matchExpression(m.accountID, asmt.Account)
if err != nil {
return false, trace.Wrap(err)
}
if !accountMatches {
continue
}
permissionSetMatches, err := matchExpression(m.permissionSetARN, asmt.PermissionSet)
if err != nil {
return false, trace.Wrap(err)
}
if permissionSetMatches {
return true, nil
}
}
return false, nil
}
func (m *IdentityCenterAccountAssignmentMatcher) String() string {
return fmt.Sprintf("IdentityCenterAccountMatcher(account=%v, permissionSet=%v)",
m.accountID, m.permissionSetARN)
}
func matchExpression(target, expression string) (bool, error) {
if expression == types.Wildcard {
return true, nil
}
matches, err := utils.MatchString(target, expression)
if err != nil {
return false, trace.Wrap(err)
}
return matches, nil
}
/*
* 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 (
"strings"
"github.com/gravitational/trace"
"github.com/vulcand/predicate"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
)
// impersonateContext is a default rule context used in teleport
type impersonateContext struct {
// user is currently authenticated user
user types.User
// impersonateRole is a role to impersonate
impersonateRole types.Role
// impersonateUser is a user to impersonate
impersonateUser types.User
}
// getIdentifier returns identifier defined in a context
func (ctx *impersonateContext) getIdentifier(fields []string) (any, error) {
switch fields[0] {
case UserIdentifier:
return predicate.GetFieldByTag(ctx.user, teleport.JSON, fields[1:])
case ImpersonateUserIdentifier:
return predicate.GetFieldByTag(ctx.impersonateUser, teleport.JSON, fields[1:])
case ImpersonateRoleIdentifier:
return predicate.GetFieldByTag(ctx.impersonateRole, teleport.JSON, fields[1:])
default:
return nil, trace.NotFound("%v is not defined", strings.Join(fields, "."))
}
}
// matchesImpersonateWhere returns true if Where rule matches.
// Empty Where block always matches.
func matchesImpersonateWhere(cond types.ImpersonateConditions, parser predicate.Parser) (bool, error) {
if cond.Where == "" {
return true, nil
}
ifn, err := parser.Parse(cond.Where)
if err != nil {
return false, trace.Wrap(err)
}
fn, ok := ifn.(predicate.BoolPredicate)
if !ok {
return false, trace.BadParameter("invalid predicate type for where expression: %v", cond.Where)
}
return fn(), nil
}
// newImpersonateWhereParser returns standard parser for `where` section in impersonate rules
func newImpersonateWhereParser(ctx *impersonateContext) (predicate.Parser, error) {
return predicate.NewParser(predicate.Def{
Operators: predicate.Operators{
AND: predicate.And,
OR: predicate.Or,
NOT: predicate.Not,
},
Functions: map[string]any{
"equals": predicate.Equals,
"contains": predicate.Contains,
},
GetIdentifier: ctx.getIdentifier,
GetProperty: GetStringMapValue,
})
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalInstaller unmarshals the installer resource from JSON.
func UnmarshalInstaller(data []byte, opts ...MarshalOption) (types.Installer, error) {
var installer types.InstallerV1
if len(data) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(data, &installer); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := installer.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
installer.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
installer.SetExpiry(cfg.Expires)
}
return &installer, nil
}
// MarshalInstaller marshals the Installer resource to JSON.
func MarshalInstaller(installer types.Installer, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch installer := installer.(type) {
case *types.InstallerV1:
if err := installer.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, installer))
default:
return nil, trace.BadParameter("unrecognized installer version %T", installer)
}
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// Integrations defines an interface for managing Integrations.
type Integrations interface {
IntegrationsGetter
// CreateIntegration creates a new integration resource.
CreateIntegration(context.Context, types.Integration) (types.Integration, error)
// UpdateIntegration updates an existing integration resource.
UpdateIntegration(context.Context, types.Integration) (types.Integration, error)
// DeleteIntegration removes the specified integration resource.
DeleteIntegration(ctx context.Context, name string) error
// DeleteAllIntegrations removes all integrations.
DeleteAllIntegrations(context.Context) error
}
// IntegrationsGetter defines methods for List/Read operations on Integration Resources.
type IntegrationsGetter interface {
// ListIntegrations returns a paginated list of all integration resources.
ListIntegrations(ctx context.Context, pageSize int, nextToken string) ([]types.Integration, string, error)
// GetIntegration returns the specified integration resources.
GetIntegration(ctx context.Context, name string) (types.Integration, error)
}
// IntegrationsTokenGenerator defines methods to generate tokens for Integrations.
type IntegrationsTokenGenerator interface {
// GenerateAWSOIDCToken generates a token to be used to execute an AWS OIDC Integration action.
GenerateAWSOIDCToken(ctx context.Context, integration string) (string, error)
// GenerateAzureOIDCToken generates a token to be used to execute an Azure OIDC Integration action.
GenerateAzureOIDCToken(ctx context.Context, integration string) (string, error)
}
// MarshalIntegration marshals the Integration resource to JSON.
func MarshalIntegration(ig types.Integration, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch g := ig.(type) {
case *types.IntegrationV1:
if err := g.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, g))
default:
return nil, trace.BadParameter("unsupported integration resource %T", g)
}
}
// UnmarshalIntegration unmarshals Integration resource from JSON.
func UnmarshalIntegration(data []byte, opts ...MarshalOption) (types.Integration, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing resource data")
}
var ig types.IntegrationV1
err := utils.FastUnmarshal(data, &ig)
if err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
ig.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
ig.SetExpiry(cfg.Expires)
}
return &ig, nil
}
/*
* Teleport
* Copyright (C) 2026 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 (
"time"
"google.golang.org/protobuf/proto"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types"
)
// KubeAccessChecker provides kube-specific access checking, abstracting over scoped and unscoped identities.
// It is obtained from [ScopedAccessChecker.Kube] and should not be constructed directly. Methods on this type
// implement kube-specific behavior, branching internally between the scoped and unscoped paths of the underlying
// [ScopedAccessChecker].
type KubeAccessChecker struct {
checker *ScopedAccessChecker
}
// CheckAccessToCluster checks access to a kube cluster.
func (c *KubeAccessChecker) CheckAccessToCluster(target types.KubeCluster, state AccessState, matchers ...RoleMatcher) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAccess(target, state, matchers...)
}
return c.checker.scopedCompatChecker.CheckAccess(target, state, matchers...)
}
// CanAccessCluster checks whether read access to the specified kube server is possible without
// regard to a specific MFA state. Used for listing/filtering.
func (c *KubeAccessChecker) CanAccessCluster(target types.KubeCluster) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAccess(target, AccessState{MFAVerified: true})
}
return c.checker.scopedCompatChecker.CheckAccess(target, AccessState{MFAVerified: true})
}
// GetGroupsAndUsers returns the kube groups and users that are permitted for impersonation.
func (c *KubeAccessChecker) GetGroupsAndUsers(ttl time.Duration, overrideTTL bool, matchers ...RoleMatcher) ([]string, []string, error) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckKubeGroupsAndUsers(ttl, overrideTTL, matchers...)
}
return c.checker.scopedCompatChecker.CheckKubeGroupsAndUsers(ttl, overrideTTL, matchers...)
}
// GetResources returns the kube resources that are permitted for access.
func (c *KubeAccessChecker) GetResources(target types.KubeCluster) (allowed []types.KubernetesResource, denied []types.KubernetesResource) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.GetKubeResources(target)
}
return c.checker.scopedCompatChecker.GetKubeResources(target)
}
// AdjustClientIdleTimeout determines the kube client idle timeout to apply. The supplied argument must be
// the globally defined most-permissive value. For scoped identities, the value is read directly from the
// scoped role proto (kube.client_idle_timeout takes precedence over defaults.client_idle_timeout). If the
// role specifies a more restrictive value it is returned; otherwise the global value is returned unchanged.
// An error is returned if the role contains a non-empty duration string that cannot be parsed.
func (c *KubeAccessChecker) AdjustClientIdleTimeout(timeout time.Duration) (time.Duration, error) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.AdjustClientIdleTimeout(timeout), nil
}
return c.checker.adjustScopedClientIdleTimeout(c.checker.role.GetSpec().GetKube().GetClientIdleTimeout(), timeout)
}
// AdjustDisconnectExpiredCert adjusts whether to disconnect on certificate expiry.
func (c *KubeAccessChecker) AdjustDisconnectExpiredCert(disconnect bool) bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.AdjustDisconnectExpiredCert(disconnect)
}
kube := c.checker.role.GetSpec().GetKube()
var disconnectExpiredCert *bool
if kube != nil {
disconnectExpiredCert = proto.ValueOrNil(kube.HasDisconnectExpiredCert(), kube.GetDisconnectExpiredCert)
}
return c.checker.adjustScopedDisconnectExpiredCert(disconnectExpiredCert, disconnect)
}
// LockingMode returns the SSH lock enforcement mode to apply.
func (c *KubeAccessChecker) LockingMode(defaultMode constants.LockingMode) constants.LockingMode {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.LockingMode(defaultMode)
}
return c.checker.scopedLockingMode(c.checker.role.GetSpec().GetKube().GetLock(), defaultMode)
}
/*
* 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"
"iter"
"github.com/gravitational/trace"
presencev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/presence/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils"
)
// KubernetesClusterGetter defines interface for fetching kubernetes cluster resources.
type KubernetesClusterGetter interface {
// GetKubernetesClusters returns all kubernetes cluster resources.
GetKubernetesClusters(context.Context) ([]types.KubeCluster, error)
// ListKubeClusters returns a page of registered kube clusters with the ability to apply
// scope filters.
ListKubeClusters(ctx context.Context, req *presencev1.ListKubeClustersRequest) ([]types.KubeCluster, string, error)
// RangeKubeClusters returns a sequence of kube clusters filtered by the given
// [*presencev1.ListKubeClustersRequest].
RangeKubeClusters(ctx context.Context, req *presencev1.ListKubeClustersRequest) iter.Seq2[types.KubeCluster, error]
// GetKubeCluster returns the specified kube cluster resource by scope and name.
GetKubeCluster(ctx context.Context, req *presencev1.GetKubeClusterRequest) (types.KubeCluster, error)
}
// KubernetesServerGetter defines interface for fetching kubernetes server resources.
type KubernetesServerGetter interface {
// GetKubernetesServers returns all kubernetes server resources.
GetKubernetesServers(context.Context) ([]types.KubeServer, error)
}
// Kubernetes defines an interface for managing kubernetes clusters resources.
type Kubernetes interface {
// KubernetesClusterGetter provides methods for fetching kubernetes resources.
KubernetesClusterGetter
// CreateKubernetesCluster creates a new kubernetes cluster resource.
CreateKubernetesCluster(context.Context, types.KubeCluster) error
// UpdateKubernetesCluster updates an existing kubernetes cluster resource.
UpdateKubernetesCluster(context.Context, types.KubeCluster) error
// DeleteAllKubernetesClusters removes all kubernetes resources.
DeleteAllKubernetesClusters(context.Context) error
// DeleteKubeCluster removes the specified kube cluster resource with
// respect to its scope.
DeleteKubeCluster(ctx context.Context, req *presencev1.DeleteKubeClusterRequest) error
}
// MarshalKubeServer marshals the KubeServer resource to JSON.
func MarshalKubeServer(kubeServer types.KubeServer, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch server := kubeServer.(type) {
case *types.KubernetesServerV3:
if err := server.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, server))
default:
return nil, trace.BadParameter("unsupported kube server resource %T", server)
}
}
// UnmarshalKubeServer unmarshals KubeServer resource from JSON.
func UnmarshalKubeServer(data []byte, opts ...MarshalOption) (types.KubeServer, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing kube server data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var s types.KubernetesServerV3
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("unsupported kube server resource version %q", h.Version)
}
// MarshalKubeCluster marshals the KubeCluster resource to JSON.
func MarshalKubeCluster(kubeCluster types.KubeCluster, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if c, ok := kubeCluster.(types.DiscoveredEKSCluster); ok {
kubeCluster = c.GetKubeCluster()
}
switch cluster := kubeCluster.(type) {
case *types.KubernetesClusterV3:
if err := cluster.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, cluster))
default:
return nil, trace.BadParameter("unsupported kube cluster resource %T", cluster)
}
}
// UnmarshalKubeCluster unmarshals KubeCluster resource from JSON.
func UnmarshalKubeCluster(data []byte, opts ...MarshalOption) (types.KubeCluster, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing kube cluster data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V3:
var s types.KubernetesClusterV3
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("unsupported kube cluster resource version %q", h.Version)
}
// GetCursorForKubeCluster returns the backend key for a kube cluster with
// consideration for whether or not it is scoped.
func GetCursorForKubeCluster(cluster types.KubeCluster) string {
return scopes.MakeResourceCursor(cluster.GetScope(), cluster.GetName())
}
// GetCursorForKubeServer returns the resource cursor identifying a kube server
// in the logical resource stream: "<host-id>/<cluster-name>" for unscoped kube servers
// and "~scoped/<encoded-scope>/<host-id>/<cluster-name>" for scoped kube servers.
func GetCursorForKubeServer(server types.KubeServer) string {
return scopes.MakeResourceCursorWithHost(server.GetScope(), server.GetHostID(), server.GetName())
}
/*
* Teleport
* Copyright (C) 2024 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/trace"
kubewaitingcontainerpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/kubewaitingcontainer/v1"
"github.com/gravitational/teleport/api/types/kubewaitingcontainer"
)
// KubeWaitingContainer is responsible for managing Kubernetes
// ephemeral containers that are waiting to be created until moderated
// session conditions are met.
type KubeWaitingContainer interface {
ListKubernetesWaitingContainers(ctx context.Context, pageSize int, pageToken string) ([]*kubewaitingcontainerpb.KubernetesWaitingContainer, string, error)
GetKubernetesWaitingContainer(ctx context.Context, req *kubewaitingcontainerpb.GetKubernetesWaitingContainerRequest) (*kubewaitingcontainerpb.KubernetesWaitingContainer, error)
CreateKubernetesWaitingContainer(ctx context.Context, in *kubewaitingcontainerpb.KubernetesWaitingContainer) (*kubewaitingcontainerpb.KubernetesWaitingContainer, error)
DeleteKubernetesWaitingContainer(ctx context.Context, req *kubewaitingcontainerpb.DeleteKubernetesWaitingContainerRequest) error
}
// MarshalKubeWaitingContainer marshals a KubernetesWaitingContainer resource to JSON.
func MarshalKubeWaitingContainer(in *kubewaitingcontainerpb.KubernetesWaitingContainer, opts ...MarshalOption) ([]byte, error) {
if err := kubewaitingcontainer.ValidateKubeWaitingContainer(in); err != nil {
return nil, trace.Wrap(err)
}
return FastMarshalProtoResourceDeprecated(in, opts...)
}
// UnmarshalKubeWaitingContainer unmarshals a KubernetesWaitingContainer resource from JSON.
func UnmarshalKubeWaitingContainer(data []byte, opts ...MarshalOption) (*kubewaitingcontainerpb.KubernetesWaitingContainer, error) {
out, err := FastUnmarshalProtoResourceDeprecated[*kubewaitingcontainerpb.KubernetesWaitingContainer](data, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
if err := kubewaitingcontainer.ValidateKubeWaitingContainer(out); err != nil {
return nil, trace.Wrap(err)
}
return out, nil
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalLicense unmarshals the License resource from JSON.
func UnmarshalLicense(bytes []byte) (types.License, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
var license types.LicenseV3
err := utils.FastUnmarshal(bytes, &license)
if err != nil {
return nil, trace.BadParameter("%s", err)
}
if license.Version != types.V3 {
return nil, trace.BadParameter("unsupported version %v, expected version %v", license.Version, types.V3)
}
if err := license.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &license, nil
}
// MarshalLicense marshals the License resource to JSON.
func MarshalLicense(license types.License, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch license := license.(type) {
case *types.LicenseV3:
if err := license.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
// avoid modifying the original object
// to prevent unexpected data races
copy := *license
copy.SetRevision("")
license = ©
}
return utils.FastMarshal(license)
default:
return nil, trace.BadParameter("unrecognized license version %T", license)
}
}
// IsDashboard returns a bool indicating if the cluster is a
// dashboard cluster.
// Dashboard is a cluster running on cloud infrastructure that
// isn't a Teleport Cloud cluster
func IsDashboard(features proto.Features) bool {
// TODO(matheus): for now, we assume dashboard based on
// the presence of recovery codes, which are never enabled
// in OSS or self-hosted Teleport.
return !features.GetCloud() && features.GetRecoveryCodes()
}
// Teleport
// Copyright (C) 2026 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"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
)
// LinuxDesktops is the interface for managing Linux desktop resources.
type LinuxDesktops interface {
LinuxDesktopGetter
// CreateLinuxDesktop creates a new Linux desktop resource.
CreateLinuxDesktop(context.Context, *linuxdesktopv1.LinuxDesktop) (*linuxdesktopv1.LinuxDesktop, error)
// UpdateLinuxDesktop updates the Linux desktop resource.
UpdateLinuxDesktop(context.Context, *linuxdesktopv1.LinuxDesktop) (*linuxdesktopv1.LinuxDesktop, error)
// UpsertLinuxDesktop updates the Linux desktop resource or create one if needed.
UpsertLinuxDesktop(context.Context, *linuxdesktopv1.LinuxDesktop) (*linuxdesktopv1.LinuxDesktop, error)
// DeleteLinuxDesktop deletes the Linux desktop resource by name.
DeleteLinuxDesktop(context.Context, string) error
}
type LinuxDesktopGetter interface {
// ListLinuxDesktops returns the Linux desktop resources.
ListLinuxDesktops(ctx context.Context, pageSize int, nextToken string) ([]*linuxdesktopv1.LinuxDesktop, string, error)
// GetLinuxDesktop returns the Linux desktop resource by name.
GetLinuxDesktop(ctx context.Context, name string) (*linuxdesktopv1.LinuxDesktop, error)
}
// MarshalLinuxDesktop marshals the LinuxDesktop object into a JSON byte array.
func MarshalLinuxDesktop(object *linuxdesktopv1.LinuxDesktop, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalLinuxDesktop unmarshals the LinuxDesktop object from a JSON byte array.
func UnmarshalLinuxDesktop(data []byte, opts ...MarshalOption) (*linuxdesktopv1.LinuxDesktop, error) {
return UnmarshalProtoResource[*linuxdesktopv1.LinuxDesktop](data, opts...)
}
/*
* 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 (
"fmt"
"iter"
"maps"
"slices"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
)
// LockInForceAccessDenied is an AccessDenied error returned when a lock
// is in force.
func LockInForceAccessDenied(lock types.Lock) error {
s := fmt.Sprintf("lock targeting %v is in force", lock.Target())
msg := lock.Message()
if len(msg) > 0 {
s += ": " + msg
}
err := trace.AccessDenied("%s", s)
return trace.WithField(err, "lock-in-force", lock)
}
// StrictLockingModeAccessDenied is an AccessDenied error returned when strict
// locking mode causes all interactions to be blocked.
var StrictLockingModeAccessDenied = trace.AccessDenied("preventive lock-out due to local lock view becoming unreliable")
// SSHAccessLockTargets computes the full set of lock targets related to ssh access.
func SSHAccessLockTargets(localClusterName, serverID, osLogin string, accessInfo *AccessInfo, unmappedIdentity *sshca.Identity) []types.LockTarget {
// ssh access lock targets are currently constructed identically to proxying lock targets,
// except that the os login associated with the ssh access attempt is also locked.
return proxyingLockTargets(localClusterName, serverID, &osLogin, accessInfo, unmappedIdentity)
}
// ProxyingLockTargets computes the full set of lock targets related to teleport proxying.
func ProxyingLockTargets(localClusterName, serverID string, accessInfo *AccessInfo, unmappedIdentity *sshca.Identity) []types.LockTarget {
var noOSLogin *string
return proxyingLockTargets(localClusterName, serverID, noOSLogin, accessInfo, unmappedIdentity)
}
func proxyingLockTargets(localClusterName, serverID string, osLogin *string, accessInfo *AccessInfo, unmappedIdentity *sshca.Identity) []types.LockTarget {
lockTargets := map[types.LockTarget]struct{}{
{User: accessInfo.Username}: struct{}{},
{ServerID: serverID}: struct{}{},
{ServerID: utils.HostFQDN(serverID, localClusterName)}: struct{}{},
}
if mfaDevice := unmappedIdentity.MFAVerified; mfaDevice != "" {
lockTargets[types.LockTarget{MFADevice: mfaDevice}] = struct{}{}
}
if trustedDevice := unmappedIdentity.DeviceID; trustedDevice != "" {
lockTargets[types.LockTarget{Device: trustedDevice}] = struct{}{}
}
if joinToken := unmappedIdentity.JoinToken; joinToken != "" {
lockTargets[types.LockTarget{JoinToken: joinToken}] = struct{}{}
}
if botInstanceID := unmappedIdentity.BotInstanceID; botInstanceID != "" {
lockTargets[types.LockTarget{BotInstanceID: botInstanceID}] = struct{}{}
}
for lockTarget := range RolesToLockTargets(slices.Values(accessInfo.Roles)) {
lockTargets[lockTarget] = struct{}{}
}
for lockTarget := range RolesToLockTargets(slices.Values(unmappedIdentity.Roles)) {
lockTargets[lockTarget] = struct{}{}
}
for lockTarget := range AccessRequestsToLockTargets(slices.Values(unmappedIdentity.ActiveRequests)) {
lockTargets[lockTarget] = struct{}{}
}
if osLogin != nil {
lockTargets[types.LockTarget{Login: *osLogin}] = struct{}{}
}
return slices.AppendSeq(make([]types.LockTarget, 0, len(lockTargets)), maps.Keys(lockTargets))
}
// GitForwardingLockTargets computes the full set of lock targets related to git forwarding.
func GitForwardingLockTargets(localClusterName, serverID string, accessInfo *AccessInfo, unmappedIdentity *sshca.Identity) []types.LockTarget {
// git forwarding lock targets are currently constructed identically to proxying lock targets.
return ProxyingLockTargets(localClusterName, serverID, accessInfo, unmappedIdentity)
}
// LockTargetsFromTLSIdentity infers a list of LockTargets from tlsca.Identity.
func LockTargetsFromTLSIdentity(id tlsca.Identity) iter.Seq[types.LockTarget] {
return func(yield func(types.LockTarget) bool) {
for lockTarget := range RolesToLockTargets(slices.Values(id.Groups)) {
if !yield(lockTarget) {
return
}
}
if !yield(types.LockTarget{User: id.Username}) {
return
}
if id.MFAVerified != "" && !yield(types.LockTarget{MFADevice: id.MFAVerified}) {
return
}
if id.DeviceExtensions.DeviceID != "" && !yield(types.LockTarget{Device: id.DeviceExtensions.DeviceID}) {
return
}
if id.JoinToken != "" && !yield(types.LockTarget{JoinToken: id.JoinToken}) {
return
}
if id.BotInstanceID != "" && !yield(types.LockTarget{BotInstanceID: id.BotInstanceID}) {
return
}
for lockTarget := range AccessRequestsToLockTargets(slices.Values(id.ActiveRequests)) {
if !yield(lockTarget) {
return
}
}
}
}
// RolesToLockTargets converts a list of roles to a list of LockTargets
// (one LockTarget per role).
func RolesToLockTargets(roles iter.Seq[string]) iter.Seq[types.LockTarget] {
return func(yield func(types.LockTarget) bool) {
for role := range roles {
if !yield(types.LockTarget{Role: role}) {
return
}
}
}
}
// AccessRequestsToLockTargets converts a list of access requests to a list of
// LockTargets (one LockTarget per access request)
func AccessRequestsToLockTargets(accessRequests iter.Seq[string]) iter.Seq[types.LockTarget] {
return func(yield func(types.LockTarget) bool) {
for accessRequest := range accessRequests {
if !yield(types.LockTarget{AccessRequest: accessRequest}) {
return
}
}
}
}
// UnmarshalLock unmarshals the Lock resource from JSON.
func UnmarshalLock(bytes []byte, opts ...MarshalOption) (types.Lock, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
var lock types.LockV2
if err := utils.FastUnmarshal(bytes, &lock); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := lock.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
lock.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
lock.SetExpiry(cfg.Expires)
}
return &lock, nil
}
// MarshalLock marshals the Lock resource to JSON.
func MarshalLock(lock types.Lock, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch lock := lock.(type) {
case *types.LockV2:
if err := lock.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if version := lock.GetVersion(); version != types.V2 {
return nil, trace.BadParameter("mismatched lock version %v and type %T", version, lock)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, lock))
default:
return nil, trace.BadParameter("unrecognized lock version %T", lock)
}
}
/*
* 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"
"log/slog"
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
azureutils "github.com/gravitational/teleport/api/utils/azure"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils"
libslices "github.com/gravitational/teleport/lib/utils/slices"
"github.com/gravitational/teleport/lib/utils/typical"
)
// ResourceMatcher matches cluster resources.
type ResourceMatcher struct {
// Labels match resource labels.
Labels types.Labels
// AWS contains AWS specific settings.
AWS ResourceMatcherAWS
}
// ResourceMatcherAWS contains AWS specific settings.
type ResourceMatcherAWS struct {
// AssumeRoleARN is the AWS role to assume for accessing the resource.
AssumeRoleARN string
// ExternalID is an optional AWS external ID used to enable assuming an AWS
// role across accounts.
ExternalID string
}
// ResourceMatchersToTypes converts []]services.ResourceMatchers into []*types.ResourceMatcher
func ResourceMatchersToTypes(in []ResourceMatcher) []*types.DatabaseResourceMatcher {
out := make([]*types.DatabaseResourceMatcher, len(in))
for i, resMatcher := range in {
out[i] = &types.DatabaseResourceMatcher{
Labels: &resMatcher.Labels,
AWS: types.ResourceMatcherAWS{
AssumeRoleARN: resMatcher.AWS.AssumeRoleARN,
ExternalID: resMatcher.AWS.ExternalID,
},
}
}
return out
}
// SimplifyAzureMatchers returns simplified Azure matchers. Each selector list
// is trimmed, deduplicated, and collapsed to the wildcard if any entry is the
// wildcard or the list is empty. Lists where every entry is whitespace-only are
// preserved verbatim rather than widened to wildcard, so invalid hand-edited
// scopes fail closed downstream instead of matching all.
func SimplifyAzureMatchers(matchers []types.AzureMatcher) []types.AzureMatcher {
result := make([]types.AzureMatcher, 0, len(matchers))
for _, m := range matchers {
subs := simplifySelector(m.Subscriptions)
groups := simplifySelector(m.ResourceGroups)
ts := apiutils.Deduplicate(m.Types)
var regions []string
if len(m.Regions) == 0 || slices.Contains(m.Regions, types.Wildcard) {
regions = []string{types.Wildcard}
} else {
// Normalize before dedup so case-variant inputs ("East US", "eastus") collapse.
// The fresh slice also keeps m.Regions safe from the watcher's concurrent IsEqual reads.
regions = apiutils.Deduplicate(libslices.Map(m.Regions, azureutils.NormalizeLocation))
}
elem := m
elem.Subscriptions = subs
elem.ResourceGroups = groups
elem.Regions = regions
elem.Types = ts
result = append(result, elem)
}
return result
}
// simplifySelector applies the SimplifyAzureMatchers normalization rules to a
// single selector list (Subscriptions or ResourceGroups). The four input
// shapes are handled distinctly to avoid silently widening malformed config
// to the wildcard:
//
// - empty input (len == 0) -> wildcard (existing convention)
// - any wildcard among trimmed entries -> wildcard (existing convention)
// - non-empty input with at least one
// non-empty trimmed entry -> trimmed, deduped entries
// - non-empty input that ALL trim to empty
// (e.g. [" "], ["", " "]) -> original input preserved verbatim
// (do NOT widen to wildcard; the Azure
// SDK call surfaces the typo)
func simplifySelector(input []string) []string {
if len(input) == 0 {
return []string{types.Wildcard}
}
trimmed := libslices.FilterMapUnique(input, utils.TrimNonEmpty)
if slices.Contains(trimmed, types.Wildcard) {
return []string{types.Wildcard}
}
if len(trimmed) == 0 {
// User supplied selectors but every entry was whitespace-only.
// This is a config typo, not a "match everything" signal.
// Preserve the original input so the Azure SDK rejects it as an
// invalid scope rather than silently broadening discovery.
return input
}
return trimmed
}
// MatchResourceLabels returns true if any of the provided selectors matches the provided database.
func MatchResourceLabels(matchers []ResourceMatcher, labels map[string]string) bool {
for _, matcher := range matchers {
if len(matcher.Labels) == 0 {
return false
}
match, _, err := MatchLabels(matcher.Labels, labels)
if err != nil {
slog.ErrorContext(context.Background(), "Failed to match labels",
"error", err,
"matcher_labels", matcher.Labels,
"resource_labels", labels,
)
return false
}
if match {
return true
}
}
return false
}
// resourceWithTargetHealth wraps a resource to provide target health info.
type resourceWithTargetHealth struct {
types.ResourceWithLabels
health types.TargetHealthStatus
}
func (r *resourceWithTargetHealth) GetTargetHealthStatus() types.TargetHealthStatus {
return r.health
}
// ResourceSeenKey is used as a key for a map that keeps track
// of unique resource names and address. Currently "addr"
// only applies to resource Application.
type ResourceSeenKey struct{ name, kind, addr, scope string }
// MatchResourcesByFilters filters provided resources with profiled filter.
func MatchResourcesByFilters[E types.ResourceWithLabels, S ~[]E](all S, filter MatchResourceFilter) (S, error) {
var filtered S
for _, r := range all {
match, err := MatchResourceByFilters(r, filter, nil)
if err != nil {
return nil, trace.Wrap(err)
} else if match {
filtered = append(filtered, r)
}
}
return filtered, nil
}
// MatchResourceByFilters returns true if all filter values given matched against the resource.
//
// If no filters were provided, we will treat that as a match.
//
// If a `seenMap` is provided, this will be treated as a request to filter out duplicate matches.
// The map will be modified in place as it adds new keys. Seen keys will return match as false.
//
// Resource KubeService is handled differently b/c of its 1-N relationhip with service-clusters,
// it filters out the non-matched clusters on the kube service and the kube service
// is modified in place with only the matched clusters. Deduplication for resource `KubeService`
// is not provided but is provided for kind `KubernetesCluster`.
func MatchResourceByFilters(resource types.ResourceWithLabels, filter MatchResourceFilter, seenMap map[ResourceSeenKey]struct{}) (bool, error) {
var specResource types.ResourceWithLabels
kind := resource.GetKind()
scope := ""
switch res := resource.(type) {
case *types.KubernetesClusterV3:
scope = res.GetScope()
case types.KubeServer:
scope = res.GetScope()
case types.AppServer:
scope = res.GetScope()
case types.Server:
scope = res.GetScope()
}
// We assume when filtering for services like KubeService, AppServer, and DatabaseServer
// the user is wanting to filter the contained resource ie. KubeClusters, Application, and Database.
key := ResourceSeenKey{
kind: kind,
name: resource.GetName(),
scope: scopes.NormalizeForEquality(scope),
}
switch kind {
case types.KindNode,
types.KindDatabaseService,
types.KindKubernetesCluster,
types.KindWindowsDesktop, types.KindWindowsDesktopService,
types.KindLinuxDesktop,
types.KindUserGroup,
types.KindIdentityCenterAccount,
types.KindIdentityCenterAccountAssignment,
types.KindGitServer:
specResource = resource
case types.KindKubeServer:
if seenMap != nil {
return false, trace.BadParameter("checking for duplicate matches for resource kind %q is not supported", filter.ResourceKind)
}
return matchAndFilterKubeClusters(resource, filter)
case types.KindDatabaseServer:
server, ok := resource.(types.DatabaseServer)
if !ok {
return false, trace.BadParameter("expected types.DatabaseServer, got %T", resource)
}
specResource = &resourceWithTargetHealth{
ResourceWithLabels: server.GetDatabase(),
health: server.GetTargetHealthStatus(),
}
key.name = specResource.GetName()
case types.KindAppServer, types.KindSAMLIdPServiceProvider:
switch appOrSP := resource.(type) {
case types.AppServer:
app := appOrSP.GetApp()
specResource = app
key.addr = app.GetPublicAddr()
key.name = app.GetName()
case types.SAMLIdPServiceProvider:
specResource = appOrSP
key.name = specResource.GetName()
default:
return false, trace.BadParameter("expected types.SAMLIdPServiceProvider or types.AppServer, got %T", resource)
}
default:
// We check if the resource kind is a Kubernetes resource kind to reduce the amount of
// of cases we need to handle. If the resource type didn't match any arm before
// and it is not a Kubernetes resource kind, we return an error.
if !slices.Contains(types.KubernetesResourcesKinds, filter.ResourceKind) && !strings.HasPrefix(filter.ResourceKind, types.AccessRequestPrefixKindKube) {
return false, trace.NotImplemented("filtering for resource kind %q not supported", kind)
}
specResource = resource
}
var match bool
if filter.IsSimple() {
match = true
}
if !match {
var err error
match, err = matchResourceByFilters(specResource, filter)
if err != nil {
return false, trace.Wrap(err)
}
}
// Deduplicate matches.
if match && seenMap != nil {
if _, exists := seenMap[key]; exists {
return false, nil
}
seenMap[key] = struct{}{}
}
return match, nil
}
func matchResourceByFilters(resource types.ResourceWithLabels, filter MatchResourceFilter) (bool, error) {
if !types.MatchKinds(resource, filter.Kinds) {
return false, nil
}
if !types.MatchLabels(resource, filter.Labels) {
return false, nil
}
if len(filter.SearchKeywords) > 0 && !resource.MatchSearch(filter.SearchKeywords) {
return false, nil
}
if filter.PredicateExpression != nil {
match, err := filter.PredicateExpression.Evaluate(resource)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
return false, nil
}
}
return true, nil
}
// matchAndFilterKubeClusters is similar to MatchResourceByFilters, but does two things in addition:
// 1. handles kube service having a 1-N relationship (service-clusters)
// so each kube cluster goes through the filters
// 2. filters out the non-matched clusters on the kube service and the kube service is
// modified in place with only the matched clusters
// 3. only returns true if the service contained any matched cluster
func matchAndFilterKubeClusters(resource types.ResourceWithLabels, filter MatchResourceFilter) (bool, error) {
if filter.IsSimple() {
return true, nil
}
switch server := resource.(type) {
case types.KubeServer:
kubeCluster := server.GetCluster()
if kubeCluster == nil {
return false, nil
}
match, err := matchResourceByFilters(&resourceWithTargetHealth{
ResourceWithLabels: kubeCluster,
health: server.GetTargetHealthStatus(),
}, filter)
return match, trace.Wrap(err)
default:
return false, trace.BadParameter("unexpected kube server of type %T", resource)
}
}
// MatchResourceFilter holds the filter values to match against a resource.
type MatchResourceFilter struct {
// ResourceKind is the resource kind and is used to fine tune the filtering.
ResourceKind string
// Labels are the labels to match.
Labels map[string]string
// SearchKeywords is a list of search keywords to match.
SearchKeywords []string
// PredicateExpression holds boolean conditions that must be matched.
PredicateExpression typical.Expression[types.ResourceWithLabels, bool]
// Kinds is a list of resourceKinds to be used when doing a unified resource query.
// It will filter out any kind not present in the list. If the list is not present or empty
// then all kinds are valid and will be returned (still subject to other included filters)
Kinds []string
}
// IsSimple is used to short-circuit matching when a filter doesn't specify anything more
// specific than resource kind.
func (m *MatchResourceFilter) IsSimple() bool {
return len(m.Labels) == 0 &&
len(m.SearchKeywords) == 0 &&
m.PredicateExpression == nil &&
len(m.Kinds) == 0
}
// MatchResourceFilterFromListResourceRequest converts a
// proto.ListResourcesRequest to MatchResourceFilter.
func MatchResourceFilterFromListResourceRequest(req *proto.ListResourcesRequest) (MatchResourceFilter, error) {
filter := MatchResourceFilter{
ResourceKind: req.ResourceType,
Labels: req.Labels,
SearchKeywords: req.SearchKeywords,
}
if req.PredicateExpression != "" {
expression, err := NewResourceExpression(req.PredicateExpression)
if err != nil {
return MatchResourceFilter{}, trace.Wrap(err)
}
filter.PredicateExpression = expression
}
return filter, nil
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// MarshalNamespace marshals the Namespace resource to JSON.
func MarshalNamespace(resource types.Namespace, opts ...MarshalOption) ([]byte, error) {
if err := resource.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, &resource))
}
// UnmarshalNamespace unmarshals the Namespace resource from JSON.
func UnmarshalNamespace(data []byte, opts ...MarshalOption) (*types.Namespace, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing namespace data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
// always skip schema validation on namespaces unmarshal
// the namespace is always created by teleport now
var namespace types.Namespace
if err := utils.FastUnmarshal(data, &namespace); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := namespace.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
namespace.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
namespace.Metadata.Expires = &cfg.Expires
}
return &namespace, nil
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalClusterNetworkingConfig unmarshals the ClusterNetworkingConfig resource from JSON.
func UnmarshalClusterNetworkingConfig(bytes []byte, opts ...MarshalOption) (types.ClusterNetworkingConfig, error) {
var netConfig types.ClusterNetworkingConfigV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &netConfig); err != nil {
return nil, trace.BadParameter("%s", err)
}
err = netConfig.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
netConfig.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
netConfig.SetExpiry(cfg.Expires)
}
return &netConfig, nil
}
// MarshalClusterNetworkingConfig marshals the ClusterNetworkingConfig resource to JSON.
func MarshalClusterNetworkingConfig(netConfig types.ClusterNetworkingConfig, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch netConfig := netConfig.(type) {
case *types.ClusterNetworkingConfigV2:
if err := netConfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, netConfig))
default:
return nil, trace.BadParameter("unrecognized cluster networking config version %T", netConfig)
}
}
/*
* Teleport
* Copyright (C) 2024 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/trace"
notificationsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/notifications/v1"
"github.com/gravitational/teleport/api/types"
)
// Notifications defines an interface for managing notifications.
type Notifications interface {
ListUserNotifications(ctx context.Context, pageSize int, startKey string) ([]*notificationsv1.Notification, string, error)
ListGlobalNotifications(ctx context.Context, pageSize int, startKey string) ([]*notificationsv1.GlobalNotification, string, error)
CreateUserNotification(ctx context.Context, notification *notificationsv1.Notification) (*notificationsv1.Notification, error)
UpsertUserNotification(ctx context.Context, notification *notificationsv1.Notification) (*notificationsv1.Notification, error)
DeleteUserNotification(ctx context.Context, username string, notificationId string) error
CreateGlobalNotification(ctx context.Context, globalNotification *notificationsv1.GlobalNotification) (*notificationsv1.GlobalNotification, error)
UpsertGlobalNotification(ctx context.Context, globalNotification *notificationsv1.GlobalNotification) (*notificationsv1.GlobalNotification, error)
DeleteGlobalNotification(ctx context.Context, notificationId string) error
UpsertUserNotificationState(ctx context.Context, username string, state *notificationsv1.UserNotificationState) (*notificationsv1.UserNotificationState, error)
DeleteUserNotificationState(ctx context.Context, username string, notificationId string) error
ListUserNotificationStates(ctx context.Context, username string, pageSize int, nextToken string) ([]*notificationsv1.UserNotificationState, string, error)
ListNotificationStatesForAllUsers(ctx context.Context, pageSize int, nextToken string) ([]*notificationsv1.UserNotificationState, string, error)
UpsertUserLastSeenNotification(ctx context.Context, username string, ulsn *notificationsv1.UserLastSeenNotification) (*notificationsv1.UserLastSeenNotification, error)
GetUserLastSeenNotification(ctx context.Context, username string) (*notificationsv1.UserLastSeenNotification, error)
DeleteUserLastSeenNotification(ctx context.Context, username string) error
// UniqueNotificationIdentifier methods should not be exposed to the client since they should only ever be used internally.
CreateUniqueNotificationIdentifier(ctx context.Context, prefix string, identifier string) (*notificationsv1.UniqueNotificationIdentifier, error)
ListUniqueNotificationIdentifiersForPrefix(ctx context.Context, prefix string, pageSize int, startKey string) ([]*notificationsv1.UniqueNotificationIdentifier, string, error)
GetUniqueNotificationIdentifier(ctx context.Context, prefix string, identifier string) (*notificationsv1.UniqueNotificationIdentifier, error)
DeleteUniqueNotificationIdentifier(ctx context.Context, prefix string, identifier string) error
}
// ValidateNotification verifies that the necessary fields are configured for a notification object.
func ValidateNotification(notification *notificationsv1.Notification) error {
if notification.GetSubKind() == "" {
return trace.BadParameter("notification subkind is missing")
}
if !notification.HasSpec() {
return trace.BadParameter("notification spec is missing")
}
if !notification.HasMetadata() {
return trace.BadParameter("notification metadata is missing")
}
if _, exists := notification.GetMetadata().GetLabels()[types.NotificationTitleLabel]; !exists {
return trace.BadParameter("notification title label is missing")
}
return nil
}
// MarshalNotification marshals a Notification resource to JSON.
func MarshalNotification(notification *notificationsv1.Notification, opts ...MarshalOption) ([]byte, error) {
if err := ValidateNotification(notification); err != nil {
return nil, trace.Wrap(err)
}
return FastMarshalProtoResourceDeprecated(notification, opts...)
}
// UnmarshalNotification unmarshals a Notification resource from JSON.
func UnmarshalNotification(data []byte, opts ...MarshalOption) (*notificationsv1.Notification, error) {
return FastUnmarshalProtoResourceDeprecated[*notificationsv1.Notification](data, opts...)
}
// ValidateGlobalNotification verifies that the necessary fields are configured for a global notification object.
func ValidateGlobalNotification(globalNotification *notificationsv1.GlobalNotification) error {
if !globalNotification.HasSpec() {
return trace.BadParameter("notification spec is missing")
}
if !globalNotification.GetSpec().HasMatcher() {
return trace.BadParameter("matcher is missing, a matcher is required for a global notification")
}
if err := ValidateNotification(globalNotification.GetSpec().GetNotification()); err != nil {
return trace.Wrap(err)
}
if globalNotification.GetSpec().GetNotification().GetSpec().GetUsername() != "" {
return trace.BadParameter("a global notification cannot have a username")
}
return nil
}
// MarshalGlobalNotification marshals a GlobalNotification resource to JSON.
func MarshalGlobalNotification(globalNotification *notificationsv1.GlobalNotification, opts ...MarshalOption) ([]byte, error) {
if err := ValidateGlobalNotification(globalNotification); err != nil {
return nil, trace.Wrap(err)
}
return MarshalProtoResource(globalNotification, opts...)
}
// UnmarshalGlobalNotification unmarshals a GlobalNotification resource from JSON.
func UnmarshalGlobalNotification(data []byte, opts ...MarshalOption) (*notificationsv1.GlobalNotification, error) {
return UnmarshalProtoResource[*notificationsv1.GlobalNotification](data, opts...)
}
// ValidateUserNotificationState verifies that the necessary fields are configured for user notification state object.
func ValidateUserNotificationState(notificationState *notificationsv1.UserNotificationState) error {
if notificationState.GetSpec().GetNotificationId() == "" {
return trace.BadParameter("notification id is missing")
}
if !notificationState.HasStatus() {
return trace.BadParameter("notification state status is missing")
}
return nil
}
// MarshalUserNotificationState marshals a UserNotificationState resource to JSON.
func MarshalUserNotificationState(notificationState *notificationsv1.UserNotificationState, opts ...MarshalOption) ([]byte, error) {
if err := ValidateUserNotificationState(notificationState); err != nil {
return nil, trace.Wrap(err)
}
return FastMarshalProtoResourceDeprecated(notificationState, opts...)
}
// UnmarshalUserNotificationState unmarshals a UserNotificationState resource from JSON.
func UnmarshalUserNotificationState(data []byte, opts ...MarshalOption) (*notificationsv1.UserNotificationState, error) {
return FastUnmarshalProtoResourceDeprecated[*notificationsv1.UserNotificationState](data, opts...)
}
// ValidateUserLastSeenNotification verifies that the necessary fields are configured for a user's last seen notification timestamp object.
func ValidateUserLastSeenNotification(lastSeenNotification *notificationsv1.UserLastSeenNotification) error {
if !lastSeenNotification.GetStatus().HasLastSeenTime() {
return trace.BadParameter("last seen time is missing")
}
return nil
}
// MarshalUserLastSeenNotification marshals a UserLastSeenNotification resource to JSON.
func MarshalUserLastSeenNotification(userLastSeenNotification *notificationsv1.UserLastSeenNotification, opts ...MarshalOption) ([]byte, error) {
if err := ValidateUserLastSeenNotification(userLastSeenNotification); err != nil {
return nil, trace.Wrap(err)
}
return FastMarshalProtoResourceDeprecated(userLastSeenNotification, opts...)
}
// UnmarshalUserLastSeenNotification unmarshals a UserLastSeenNotification resource from JSON.
func UnmarshalUserLastSeenNotification(data []byte, opts ...MarshalOption) (*notificationsv1.UserLastSeenNotification, error) {
return FastUnmarshalProtoResourceDeprecated[*notificationsv1.UserLastSeenNotification](data, opts...)
}
// ValidateUniqueNotificationIdentifier verifies that the necessary fields are configured for a unique notification identifier object.
func ValidateUniqueNotificationIdentifier(uni *notificationsv1.UniqueNotificationIdentifier) error {
if uni.GetSpec().GetUniqueIdentifier() == "" {
return trace.BadParameter("unique notification identifier key is missing")
}
return nil
}
// MarshalUniqueNotificationIdentifier marshals a UniqueNotificationIdentifier resource to JSON.
func MarshalUniqueNotificationIdentifier(uni *notificationsv1.UniqueNotificationIdentifier, opts ...MarshalOption) ([]byte, error) {
if err := ValidateUniqueNotificationIdentifier(uni); err != nil {
return nil, trace.Wrap(err)
}
return FastMarshalProtoResourceDeprecated(uni, opts...)
}
// UnmarshalUniqueNotificationIdentifier unmarshals a UniqueNotificationIdentifier resource from JSON.
func UnmarshalUniqueNotificationIdentifier(data []byte, opts ...MarshalOption) (*notificationsv1.UniqueNotificationIdentifier, error) {
return FastUnmarshalProtoResourceDeprecated[*notificationsv1.UniqueNotificationIdentifier](data, opts...)
}
/*
* Teleport
* Copyright (C) 2024 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"
"fmt"
"io"
"log/slog"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
notificationsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/notifications/v1"
"github.com/gravitational/teleport/api/internalutils/stream"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/sortcache"
)
type notificationsCacheIndex string
const (
// notificationKey is the key for a user-specific notification in the format of <username>/<notification uuid>.
// This index is only used by the user notifications cache. Since UUIDv7's contain a timestamp and are lexicographically sortable
// by date, this is what will be used to sort by date.
notificationKey notificationsCacheIndex = "Key"
// notificationID is the uuid of a notification.
notificationID notificationsCacheIndex = "ID"
)
// NotificationGetter defines the interface for fetching notifications.
type NotificationGetter interface {
// ListUserNotifications returns a paginated list of user-specific notifications for all users.
ListUserNotifications(ctx context.Context, pageSize int, startKey string) ([]*notificationsv1.Notification, string, error)
// ListGlobalNotifications returns a paginated list of global notifications.
ListGlobalNotifications(ctx context.Context, pageSize int, startKey string) ([]*notificationsv1.GlobalNotification, string, error)
}
// UserNotificationsCacheConfig holds the configuration parameters for both [UserNotificationCache] and [GlobalNotificationCache].
type NotificationCacheConfig struct {
// Clock is a clock for time-related operation.
Clock clockwork.Clock
// Events is an event system client.
Events types.Events
// Getter is an notification getter client.
Getter NotificationGetter
}
// CheckAndSetDefaults validates the config and provides reasonable defaults for optional fields.
func (c *NotificationCacheConfig) CheckAndSetDefaults() error {
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
if c.Events == nil {
return trace.BadParameter("notification cache config missing event system client")
}
if c.Getter == nil {
return trace.BadParameter("notification cache config missing notifications getter")
}
return nil
}
// UserNotificationCache is a custom cache for user-specific notifications, this is to allow
// fetching notifications by date in descending order.
type UserNotificationCache struct {
rw sync.RWMutex
cfg NotificationCacheConfig
primaryCache *sortcache.SortCache[*notificationsv1.Notification, notificationsCacheIndex]
ttlCache *utils.FnCache
initC chan struct{}
closeContext context.Context
cancel context.CancelFunc
}
// GlobalNotificationCache is a custom cache for user-specific notifications, this is to allow
// fetching notifications by date in descending order.
type GlobalNotificationCache struct {
rw sync.RWMutex
cfg NotificationCacheConfig
primaryCache *sortcache.SortCache[*notificationsv1.GlobalNotification, notificationsCacheIndex]
ttlCache *utils.FnCache
initC chan struct{}
closeContext context.Context
cancel context.CancelFunc
}
// NewUserNotificationCache sets up a new [UserNotificationCache] instance based on the supplied
// configuration. The cache is initialized asychronously in the background, so while it is
// safe to read from it immediately, performance is better after the cache properly initializes.
func NewUserNotificationCache(cfg NotificationCacheConfig) (*UserNotificationCache, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
ctx, cancel := context.WithCancel(context.Background())
ttlCache, err := utils.NewFnCache(utils.FnCacheConfig{
Context: ctx,
TTL: 15 * time.Second,
Clock: cfg.Clock,
})
if err != nil {
cancel()
return nil, trace.Wrap(err)
}
c := &UserNotificationCache{
cfg: cfg,
ttlCache: ttlCache,
initC: make(chan struct{}),
closeContext: ctx,
cancel: cancel,
}
if _, err := newResourceWatcher(ctx, c, ResourceWatcherConfig{
Component: "user-notification-cache",
Client: cfg.Events,
}); err != nil {
cancel()
return nil, trace.Wrap(err)
}
return c, nil
}
// NewGlobalNotificationCache sets up a new [GlobalNotificationCache] instance based on the supplied
// configuration. The cache is initialized asychronously in the background, so while it is
// safe to read from it immediately, performance is better after the cache properly initializes.
func NewGlobalNotificationCache(cfg NotificationCacheConfig) (*GlobalNotificationCache, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
ctx, cancel := context.WithCancel(context.Background())
ttlCache, err := utils.NewFnCache(utils.FnCacheConfig{
Context: ctx,
TTL: 15 * time.Second,
Clock: cfg.Clock,
})
if err != nil {
cancel()
return nil, trace.Wrap(err)
}
c := &GlobalNotificationCache{
cfg: cfg,
ttlCache: ttlCache,
initC: make(chan struct{}),
closeContext: ctx,
cancel: cancel,
}
if _, err := newResourceWatcher(ctx, c, ResourceWatcherConfig{
Component: "global-notification-cache",
Client: cfg.Events,
}); err != nil {
cancel()
return nil, trace.Wrap(err)
}
return c, nil
}
// StreamUserNotifications returns a stream with the user-specific notifications in the cache for a specified user, sorted from newest to oldest.
// We use streams here as it's a convenient way for us to construct pages to be returned to the UI one item at a time in combination with global notifications.
func (c *UserNotificationCache) StreamUserNotifications(ctx context.Context, username string, startKey string) stream.Stream[*notificationsv1.Notification] {
if username == "" {
return stream.Fail[*notificationsv1.Notification](trace.BadParameter("username is required for fetching user notifications"))
}
endKey := username + string(backend.Separator)
// Get the initial startKey if it wasn't provided.
if startKey == "" {
startKey = sortcache.NextKey(endKey)
} else {
// The sortcache expects the key to be in <username>/<uuid> format, so we prepend the username since the startKey passed into this function will just be a UUID.
startKey = username + "/" + startKey
}
const limit = 50
var done bool
return stream.PageFunc(func() ([]*notificationsv1.Notification, error) {
if done {
return nil, io.EOF
}
cache, err := c.read(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
if !cache.HasIndex(notificationKey) {
return nil, trace.Errorf("user notifications cache was not configured with index \"" + string(notificationKey) + "\" (this is a bug)")
}
notifications := make([]*notificationsv1.Notification, 0, limit)
for n := range cache.Descend(notificationKey, startKey, endKey) {
if len(notifications) == limit {
startKey = cache.KeyOf(notificationKey, n)
return notifications, nil
}
notifications = append(notifications, apiutils.CloneProtoMsg(n))
}
done = true
return notifications, nil
})
}
// fetch initializes a sortcache with all existing user-specific notifications. This is used to set up the initialize the primary cache, and
// to create a temporary cache as a fallback in case the primary cache is ever unhealthy.
func (c *UserNotificationCache) fetch(ctx context.Context) (*sortcache.SortCache[*notificationsv1.Notification, notificationsCacheIndex], error) {
cache := sortcache.New(sortcache.Config[*notificationsv1.Notification, notificationsCacheIndex]{
Indexes: map[notificationsCacheIndex]func(*notificationsv1.Notification) string{
notificationKey: func(n *notificationsv1.Notification) string {
return GetUserSpecificKey(n)
},
notificationID: func(n *notificationsv1.Notification) string {
return n.GetMetadata().GetName()
},
},
})
var startKey string
for {
notifications, nextKey, err := c.cfg.Getter.ListUserNotifications(ctx, 0, startKey)
if err != nil {
return nil, trace.Wrap(err)
}
for _, n := range notifications {
if evicted := cache.Put(n); evicted != 0 {
// this warning, if it appears, means that we configured our indexes incorrectly and one notification is overwriting another.
// the most likely explanation is that one of our indexes is missing the notification id suffix we typically use.
slog.WarnContext(ctx, "Notification conflicted with other notifications during cache fetch. This is a bug and may result in notifications not appearing the in UI correctly.", "notification", n.GetMetadata().GetName(), "num_clashing_notifications", evicted)
}
}
if nextKey == "" {
break
}
startKey = nextKey
}
return cache, nil
}
// GetUserSpecificKey returns the key for a user-specific notification in <username>/<notification uuid> format.
func GetUserSpecificKey(n *notificationsv1.Notification) string {
username := n.GetSpec().GetUsername()
id := n.GetMetadata().GetName()
return fmt.Sprintf("%s/%s", username, id)
}
// read gets a read-only view into a valid cache state. it prefers reading from the primary cache, but will fallback
// to a periodically reloaded temporary state when the primary state is unhealthy.
func (c *UserNotificationCache) read(ctx context.Context) (*sortcache.SortCache[*notificationsv1.Notification, notificationsCacheIndex], error) {
c.rw.RLock()
primary := c.primaryCache
c.rw.RUnlock()
// primary cache state is healthy, so use that. note that we don't protect access to the sortcache itself
// via our rw lock. sortcaches have their own internal locking. we just use our lock to protect the *pointer*
// to the sortcache.
if primary != nil {
return primary, nil
}
temp, err := utils.FnCacheGet(ctx, c.ttlCache, "user-notification-cache", func(ctx context.Context) (*sortcache.SortCache[*notificationsv1.Notification, notificationsCacheIndex], error) {
return c.fetch(ctx)
})
// primary may have been concurrently loaded. if it was, prefer using that.
c.rw.RLock()
primary = c.primaryCache
c.rw.RUnlock()
if primary != nil {
return primary, nil
}
return temp, trace.Wrap(err)
}
// StreamGlobalNotifications returns a stream with all the global notifications in the cache, sorted from newest to oldest.
func (c *GlobalNotificationCache) StreamGlobalNotifications(ctx context.Context, startKey string) stream.Stream[*notificationsv1.GlobalNotification] {
const limit = 50
var done bool
return stream.PageFunc(func() ([]*notificationsv1.GlobalNotification, error) {
if done {
return nil, io.EOF
}
cache, err := c.read(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
if !cache.HasIndex(notificationID) {
return nil, trace.Errorf("global notifications cache was not configured with index \"" + string(notificationID) + "\" (this is a bug)")
}
notifications := make([]*notificationsv1.GlobalNotification, 0, limit)
for n := range cache.Descend(notificationID, startKey, "") {
if len(notifications) == limit {
startKey = cache.KeyOf(notificationID, n)
return notifications, nil
}
notifications = append(notifications, apiutils.CloneProtoMsg(n))
}
done = true
return notifications, nil
})
}
// fetch initializes a sortcache with all existing global notifications. This is used to set up the initialize the primary cache, and
// to create a temporary cache as a fallback in case the primary cache is ever unhealthy.
func (c *GlobalNotificationCache) fetch(ctx context.Context) (*sortcache.SortCache[*notificationsv1.GlobalNotification, notificationsCacheIndex], error) {
cache := sortcache.New(sortcache.Config[*notificationsv1.GlobalNotification, notificationsCacheIndex]{
Indexes: map[notificationsCacheIndex]func(*notificationsv1.GlobalNotification) string{
notificationID: func(gn *notificationsv1.GlobalNotification) string {
return gn.GetMetadata().GetName()
},
},
})
var startKey string
for {
notifications, nextKey, err := c.cfg.Getter.ListGlobalNotifications(ctx, 0, startKey)
if err != nil {
return nil, trace.Wrap(err)
}
for _, n := range notifications {
if evicted := cache.Put(n); evicted != 0 {
// this warning, if it appears, means that we configured our indexes incorrectly and one notification is overwriting another.
// the most likely explanation is that one of our indexes is missing the notification id suffix we typically use.
slog.WarnContext(ctx, "Notification conflicted with other notifications during cache fetch. This is a bug and may result in notifications not appearing the in UI correctly.", "notification", n.GetMetadata().GetName(), "num_clashing_notifications", evicted)
}
}
if nextKey == "" {
break
}
startKey = nextKey
}
return cache, nil
}
// read gets a read-only view into a valid cache state. it prefers reading from the primary cache, but will fallback
// to a periodically reloaded temporary state when the primary state is unhealthy.
func (c *GlobalNotificationCache) read(ctx context.Context) (*sortcache.SortCache[*notificationsv1.GlobalNotification, notificationsCacheIndex], error) {
c.rw.RLock()
primary := c.primaryCache
c.rw.RUnlock()
// primary cache state is healthy, so use that. note that we don't protect access to the sortcache itself
// via our rw lock. sortcaches have their own internal locking. we just use our lock to protect the *pointer*
// to the sortcache.
if primary != nil {
return primary, nil
}
temp, err := utils.FnCacheGet(ctx, c.ttlCache, "global-notification-cache", func(ctx context.Context) (*sortcache.SortCache[*notificationsv1.GlobalNotification, notificationsCacheIndex], error) {
return c.fetch(ctx)
})
// primary may have been concurrently loaded. if it was, prefer using that.
c.rw.RLock()
primary = c.primaryCache
c.rw.RUnlock()
if primary != nil {
return primary, nil
}
return temp, trace.Wrap(err)
}
// --- the below methods implement the resourceCollector interface ---
// resourceKinds is part of the resourceCollector interface and is used to configure the event watcher
// that monitors for notification modifications.
func (c *UserNotificationCache) resourceKinds() []types.WatchKind {
return []types.WatchKind{
{
Kind: types.KindNotification,
},
}
}
func (c *GlobalNotificationCache) resourceKinds() []types.WatchKind {
return []types.WatchKind{
{
Kind: types.KindGlobalNotification,
},
}
}
// getResourcesAndUpdateCurrent is part of the resourceCollector interface and is called when the
// event stream for the cache has been initialized to trigger setup of the initial primary cache state.
func (c *UserNotificationCache) getResourcesAndUpdateCurrent(ctx context.Context) error {
cache, err := c.fetch(ctx)
if err != nil {
return trace.Wrap(err)
}
c.rw.Lock()
defer c.rw.Unlock()
c.primaryCache = cache
return nil
}
func (c *GlobalNotificationCache) getResourcesAndUpdateCurrent(ctx context.Context) error {
cache, err := c.fetch(ctx)
if err != nil {
return trace.Wrap(err)
}
c.rw.Lock()
defer c.rw.Unlock()
c.primaryCache = cache
return nil
}
// processEventsAndUpdateCurrent is part of the resourceCollector interface and is used to update the
// primary cache state when modification events occur.
func (c *UserNotificationCache) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
if len(events) < 1 {
return
}
c.rw.RLock()
cache := c.primaryCache
c.rw.RUnlock()
if cache == nil {
return
}
for _, event := range events {
switch event.Type {
case types.OpPut:
// Since the EventsService watcher currently only supports legacy resources, we had to use types.Resource153ToLegacy() when parsing the event
// to transform the notification into a legacy resource. We now have to use Unwrap() to get the original RFD153-style notification out and add it to the cache.
resource153, ok := event.Resource.(types.Resource153UnwrapperT[*notificationsv1.Notification])
if !ok {
slog.WarnContext(ctx, "Unexpected resource type in event (expected types.Resource153Unwrapper)", "resource_type", logutils.TypeAttr(resource153))
continue
}
notification := resource153.UnwrapT()
if evicted := cache.Put(notification); evicted > 1 {
slog.WarnContext(ctx, "Processing of put event for notification resulted in multiple cache evictions (this is a bug).", "notification", notification.GetMetadata().GetName())
}
case types.OpDelete:
cache.Delete(notificationID, event.Resource.GetName())
default:
slog.WarnContext(ctx, "Unexpected event variant", "event", event.Type)
}
}
}
func (c *GlobalNotificationCache) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
if len(events) < 1 {
return
}
c.rw.RLock()
cache := c.primaryCache
c.rw.RUnlock()
if cache == nil {
return
}
for _, event := range events {
switch event.Type {
case types.OpPut:
resource153, ok := event.Resource.(types.Resource153UnwrapperT[*notificationsv1.GlobalNotification])
if !ok {
slog.WarnContext(ctx, "Unexpected resource type in event (expected types.Resource153Unwrapper)", "resource_type", logutils.TypeAttr(resource153))
continue
}
globalNotification := resource153.UnwrapT()
if evicted := cache.Put(globalNotification); evicted > 1 {
slog.WarnContext(ctx, "Processing of put event for notification resulted in multiple cache evictions (this is a bug).", "notification", globalNotification.GetMetadata().GetName())
}
case types.OpDelete:
cache.Delete(notificationID, event.Resource.GetName())
default:
slog.WarnContext(ctx, "Unexpected event variant", "event", event.Type)
}
}
}
// notifyStale is part of the resourceCollector interface and is used to inform
// the notification cache that its view is outdated (presumably due to issues with
// the event stream).
func (c *UserNotificationCache) notifyStale() {
c.rw.Lock()
defer c.rw.Unlock()
if c.primaryCache == nil {
return
}
c.primaryCache = nil
c.initC = make(chan struct{})
}
func (c *GlobalNotificationCache) notifyStale() {
c.rw.Lock()
defer c.rw.Unlock()
if c.primaryCache == nil {
return
}
c.primaryCache = nil
c.initC = make(chan struct{})
}
// initializationChan is part of the resourceCollector interface and gets the channel
// used to signal that the notification cache has been initialized.
func (c *UserNotificationCache) initializationChan() <-chan struct{} {
c.rw.RLock()
defer c.rw.RUnlock()
return c.initC
}
func (c *GlobalNotificationCache) initializationChan() <-chan struct{} {
c.rw.RLock()
defer c.rw.RUnlock()
return c.initC
}
// Close terminates the background process that keeps the notification cache up to
// date, and terminates any inflight load operations.
func (c *UserNotificationCache) Close() error {
c.cancel()
return nil
}
func (c *GlobalNotificationCache) Close() error {
c.cancel()
return nil
}
/*
* 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 (
"net/url"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// GetRedirectURL gets a redirect URL for the given connector. If the connector
// has a redirect URL which matches the host of the given Proxy address, then
// that one will be returned. Otherwise, the first URL in the list will be returned.
func GetRedirectURL(conn types.OIDCConnector, proxyAddr string) (string, error) {
if len(conn.GetRedirectURLs()) == 0 {
return "", trace.BadParameter("No redirect URLs provided")
}
// If a specific proxyAddr wasn't provided in the oidc auth request,
// or there is only one redirect URL, use the first redirect URL.
if proxyAddr == "" || len(conn.GetRedirectURLs()) == 1 {
return conn.GetRedirectURLs()[0], nil
}
proxyNetAddr, err := utils.ParseAddr(proxyAddr)
if err != nil {
return "", trace.Wrap(err, "invalid proxy address %v", proxyAddr)
}
var matchingHostname string
for _, r := range conn.GetRedirectURLs() {
redirectURL, err := url.ParseRequestURI(r)
if err != nil {
return "", trace.Wrap(err)
}
// If we have a direct host:port match, return it.
if proxyNetAddr.String() == redirectURL.Host {
return r, nil
}
// If we have a matching host, but not port,
// save it as the best match for now.
if matchingHostname == "" && proxyNetAddr.Host() == redirectURL.Hostname() {
matchingHostname = r
}
}
if matchingHostname != "" {
return matchingHostname, nil
}
// No match, default to the first redirect URL.
return conn.GetRedirectURLs()[0], nil
}
// UnmarshalOIDCConnector unmarshals the OIDCConnector resource from JSON.
func UnmarshalOIDCConnector(bytes []byte, opts ...MarshalOption) (types.OIDCConnector, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
err = utils.FastUnmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
// V2 and V3 have the same layout, the only change is in the behavior
case types.V2, types.V3:
var c types.OIDCConnectorV3
if err := utils.FastUnmarshal(bytes, &c); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := c.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
c.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
c.SetExpiry(cfg.Expires)
}
return &c, nil
}
return nil, trace.BadParameter("OIDC connector resource version %v is not supported", h.Version)
}
// MarshalOIDCConnector marshals the OIDCConnector resource to JSON.
func MarshalOIDCConnector(oidcConnector types.OIDCConnector, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch oidcConnector := oidcConnector.(type) {
case *types.OIDCConnectorV3:
if err := oidcConnector.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, oidcConnector))
default:
return nil, trace.BadParameter("unrecognized OIDC connector version %T", oidcConnector)
}
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/client/okta"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// Compile time checks for the Okta client.
var _ OktaImportRules = (*okta.Client)(nil)
var _ OktaAssignments = (*okta.Client)(nil)
// Okta is an Okta interface for both the rules and assignments.
type Okta interface {
OktaImportRules
OktaAssignments
}
// OktaImportRules defines an interface for managing OktaImportRules.
type OktaImportRules interface {
// ListOktaImportRules returns a paginated list of all Okta import rule resources.
ListOktaImportRules(context.Context, int, string) ([]types.OktaImportRule, string, error)
// GetOktaImportRule returns the specified Okta import rule resources.
GetOktaImportRule(ctx context.Context, name string) (types.OktaImportRule, error)
// CreateOktaImportRule creates a new Okta import rule resource.
CreateOktaImportRule(context.Context, types.OktaImportRule) (types.OktaImportRule, error)
// UpdateOktaImportRule updates an existing Okta import rule resource.
UpdateOktaImportRule(context.Context, types.OktaImportRule) (types.OktaImportRule, error)
// DeleteOktaImportRule removes the specified Okta import rule resource.
DeleteOktaImportRule(ctx context.Context, name string) error
// DeleteAllOktaImportRules removes all Okta import rules.
DeleteAllOktaImportRules(context.Context) error
}
// OktaAssignmentsGetter defines an interface for reading OktaAssignments.
type OktaAssignmentsGetter interface {
// ListOktaAssignments returns a paginated list of all Okta assignment resources.
ListOktaAssignments(context.Context, int, string) ([]types.OktaAssignment, string, error)
// GetOktaAssignment returns the specified Okta assignment resources.
GetOktaAssignment(ctx context.Context, name string) (types.OktaAssignment, error)
}
// OktaAssignments defines an interface for managing OktaAssignments.
type OktaAssignments interface {
OktaAssignmentsGetter
// CreateOktaAssignment creates a new Okta assignment resource.
CreateOktaAssignment(context.Context, types.OktaAssignment) (types.OktaAssignment, error)
// UpdateOktaAssignment updates an existing Okta assignment resource.
UpdateOktaAssignment(context.Context, types.OktaAssignment) (types.OktaAssignment, error)
// ConditionalUpdateOktaAssignment updates an existing Okta assignment resource, protected by optimistic locking.
ConditionalUpdateOktaAssignment(context.Context, types.OktaAssignment) (types.OktaAssignment, error)
// UpsertOktaAssignment upsert an Okta assignment.
UpsertOktaAssignment(context.Context, types.OktaAssignment) (types.OktaAssignment, error)
// DeleteOktaAssignment removes the specified Okta assignment resource.
DeleteOktaAssignment(ctx context.Context, name string) error
// ConditionalDeleteOktaAssignment removes the specified Okta assignment resource, protected by optimistic locking.
ConditionalDeleteOktaAssignment(ctx context.Context, name, revision string) error
// DeleteAllOktaAssignments removes all Okta assignments.
DeleteAllOktaAssignments(context.Context) error
}
// MarshalOktaImportRule marshals the Okta import rule resource to JSON.
func MarshalOktaImportRule(importRule types.OktaImportRule, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch i := importRule.(type) {
case *types.OktaImportRuleV1:
if err := i.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, i))
default:
return nil, trace.BadParameter("unsupported Okta import rule resource %T", i)
}
}
// UnmarshalOktaImportRule unmarshals Okta import rule resource from JSON.
func UnmarshalOktaImportRule(data []byte, opts ...MarshalOption) (types.OktaImportRule, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing Okta import rule data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var i types.OktaImportRuleV1
if err := utils.FastUnmarshal(data, &i); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := i.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
i.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
i.SetExpiry(cfg.Expires)
}
return &i, nil
}
return nil, trace.BadParameter("unsupported Okta import rule resource version %q", h.Version)
}
// MarshalOktaAssignment marshals the Okta assignment resource to JSON.
func MarshalOktaAssignment(assignment types.OktaAssignment, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch a := assignment.(type) {
case *types.OktaAssignmentV1:
if err := a.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, a))
default:
return nil, trace.BadParameter("unsupported Okta assignment resource %T", a)
}
}
// UnmarshalOktaAssignment unmarshals the Okta assignment resource from JSON.
func UnmarshalOktaAssignment(data []byte, opts ...MarshalOption) (types.OktaAssignment, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing Okta assignment data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var a types.OktaAssignmentV1
if err := utils.FastUnmarshal(data, &a); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := a.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
a.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
a.SetExpiry(cfg.Expires)
}
return &a, nil
}
return nil, trace.BadParameter("unsupported Okta assignment resource version %q", h.Version)
}
/*
* 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"
"fmt"
"go/ast"
goparser "go/parser"
"go/token"
"log/slog"
"strconv"
"strings"
"sync"
"time"
"unicode/utf8"
"github.com/gravitational/trace"
"github.com/vulcand/predicate"
"github.com/vulcand/predicate/builder"
"golang.org/x/tools/go/ast/astutil"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/api/types/wrappers"
"github.com/gravitational/teleport/lib/session"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/set"
"github.com/gravitational/teleport/lib/utils/typical"
)
// RuleContext specifies context passed to the
// rule processing matcher, and contains information
// about current session, e.g. current user
type RuleContext interface {
// GetIdentifier returns identifier defined in a context
GetIdentifier(fields []string) (any, error)
// GetResource returns resource if specified in the context,
// if unspecified, returns error.
GetResource() (types.Resource, error)
// GetAccessChecker returns access checker if specified in the context,
// if unspecified, returns error.
GetAccessChecker() (AccessChecker, error)
}
var (
// ResourceNameExpr is the identifier that specifies resource name.
ResourceNameExpr = builder.Identifier("resource.metadata.name")
// CertAuthorityTypeExpr is a function call that returns
// cert authority type.
CertAuthorityTypeExpr = builder.Identifier(`system.catype()`)
)
// predicateAllEndWith is a custom function to test if a string ends with a
// particular suffix. If given a `[]string` as the first argument, all values
// must have the given suffix (2nd argument).
func predicateAllEndWith(a any, b any) predicate.BoolPredicate {
return func() bool {
// bval is the suffix and must always be a plain string.
bval, ok := b.(string)
if !ok {
return false
}
switch aval := a.(type) {
case string:
return strings.HasSuffix(aval, bval)
case []string:
for _, val := range aval {
if !strings.HasSuffix(val, bval) {
return false
}
}
return true
default:
return false
}
}
}
// predicateAllEqual is a custom function to test if all entries in a []string
// are equal to a certain value. This is primarily useful for comparing string
// fields that are only expected to contain a single, specific value.
func predicateAllEqual(a any, b any) predicate.BoolPredicate {
return func() bool {
// bval is the suffix and must always be a plain string.
bval, ok := b.(string)
if !ok {
return false
}
switch aval := a.(type) {
case string:
return aval == bval
case []string:
for _, val := range aval {
if val != bval {
return false
}
}
return true
default:
return false
}
}
}
// predicateIsSubset determines if the first parameter is contained within the
// variadic args. The first argument may either by `string` or `[]string`, and
// the variadic args may only be `string`.
func predicateIsSubset(a any, b ...any) predicate.BoolPredicate {
return func() bool {
// Populate the set.
set := map[string]bool{}
for _, bval := range b {
s, ok := bval.(string)
if !ok {
return false
}
set[s] = true
}
switch aval := a.(type) {
case string:
return set[aval]
case []string:
for _, v := range aval {
if !set[v] {
return false
}
}
return true
default:
return false
}
}
}
func newDefaultWhereParserDef(ctx RuleContext) predicate.Def {
def := predicate.Def{
Operators: predicate.Operators{
AND: predicate.And,
OR: predicate.Or,
NOT: predicate.Not,
EQ: predicate.Equals,
NEQ: func(a any, b any) predicate.BoolPredicate {
return func() bool {
return !predicate.Equals(a, b)()
}
},
},
Functions: map[string]any{
"equals": predicate.Equals,
"contains": predicate.Contains,
"contains_all": predicateContainsAll,
"contains_any": predicateContainsAny,
"set": func(a ...any) []string {
aVal := make([]string, 0, len(a))
for _, v := range a {
if str, ok := v.(string); ok {
aVal = append(aVal, str)
}
}
return aVal
},
"all_end_with": predicateAllEndWith,
"all_equal": predicateAllEqual,
"is_subset": predicateIsSubset,
// system.catype is a function that returns cert authority type,
// it returns empty values for unrecognized values to
// pass static rule checks.
"system.catype": func() (any, error) {
resource, err := ctx.GetResource()
if err != nil {
if trace.IsNotFound(err) {
return "", nil
}
return nil, trace.Wrap(err)
}
ca, ok := resource.(types.CertAuthority)
if !ok {
return "", nil
}
return string(ca.GetType()), nil
},
"has_prefix": func(a, b any) predicate.BoolPredicate {
return func() bool {
aval, ok := a.(string)
if !ok {
return false
}
bval, ok := b.(string)
if !ok {
return false
}
return strings.HasPrefix(aval, bval)
}
},
},
GetIdentifier: ctx.GetIdentifier,
GetProperty: GetStringMapValue,
}
return def
}
// WhereParserOpt is a function that modifies the default
// predicate.Def used to create a parser for the `where` section in access rules.
type WhereParserOpt func(RuleContext, *predicate.Def)
// WithCanViewFunction adds a can_view function to the parser definition.
// This function will be used to check if the user has access to the resource
// specified in the context.
func WithCanViewFunction() WhereParserOpt {
return func(ctx RuleContext, def *predicate.Def) {
def.Functions["can_view"] = CanViewResourceFunc(ctx)
}
}
// ConditionalOption applies the specified option if the condition is true.
// Otherwise, the returned option is a no-op.
func ConditionalOption(condition bool, option WhereParserOpt) WhereParserOpt {
return func(ctx RuleContext, def *predicate.Def) {
if condition {
option(ctx, def)
}
}
}
// NewWhereParser returns standard parser for `where` section in access rules.
func NewWhereParser(ctx RuleContext, opts ...WhereParserOpt) (predicate.Parser, error) {
def := newDefaultWhereParserDef(ctx)
for _, opt := range opts {
opt(ctx, &def)
}
return predicate.NewParser(def)
}
// CanViewResourceFunc returns a function that checks if the user has access
// to the resource specified in the context.
func CanViewResourceFunc(ctx RuleContext) func() predicate.BoolPredicate {
return func() predicate.BoolPredicate {
return func() bool {
resource, err := ctx.GetResource()
if err != nil {
return false
}
accessCheckableResource, ok := resource.(AccessCheckable)
if !ok {
return false
}
checker, err := ctx.GetAccessChecker()
if err != nil {
return false
}
// We do not enforce MFA or Device Trust for this check because
// we don't have a way of checking it from the context.
return checker.CheckAccess(accessCheckableResource, AccessState{
MFARequired: MFARequiredNever,
MFAVerified: true,
}) == nil
}
}
}
// GetStringMapValue is a helper function that returns property
// from map[string]string or map[string][]string
// the function returns empty value in case if key not found
// In case if map is nil, returns empty value as well
func GetStringMapValue(mapVal, keyVal any) (any, error) {
key, ok := keyVal.(string)
if !ok {
return nil, trace.BadParameter("only string keys are supported")
}
switch m := mapVal.(type) {
case map[string][]string:
if len(m) == 0 {
// to return nil with a proper type
var n []string
return n, nil
}
return m[key], nil
case wrappers.Traits:
if len(m) == 0 {
// to return nil with a proper type
var n []string
return n, nil
}
return m[key], nil
case map[string]string:
if len(m) == 0 {
return "", nil
}
return m[key], nil
default:
_, ok := mapVal.(map[string][]string)
return nil, trace.BadParameter("type %T is not supported, but %v %#v", m, ok, mapVal)
}
}
// NewActionsParser returns standard parser for 'actions' section in access rules
func NewActionsParser(ctx RuleContext) (predicate.Parser, error) {
return predicate.NewParser(predicate.Def{
Operators: predicate.Operators{},
Functions: map[string]any{
"log": NewLogActionFn(ctx),
},
GetIdentifier: ctx.GetIdentifier,
GetProperty: predicate.GetStringMapValue,
})
}
// NewLogActionFn creates logger functions
func NewLogActionFn(ctx RuleContext) any {
l := &LogAction{ctx: ctx}
return l.Log
}
// LogAction represents action that will emit log entry
// when specified in the actions of a matched rule
type LogAction struct {
ctx RuleContext
}
// Log logs with specified level and message string and attributes
func (l *LogAction) Log(level, msg string, args ...any) predicate.BoolPredicate {
return func() bool {
slevel := slog.LevelDebug
switch strings.ToLower(level) {
case "error":
slevel = slog.LevelError
case "warn", "warning":
slevel = slog.LevelWarn
case "info":
slevel = slog.LevelInfo
case "debug":
slevel = slog.LevelDebug
case "trace":
slevel = logutils.TraceLevel
}
ctx := context.Background()
// Expicitly check whether logging is enabled for the level
// to avoid formatting the message if the log won't be sampled.
if slog.Default().Handler().Enabled(ctx, slevel) {
//nolint:sloglint // msg cannot be constant
slog.Log(context.Background(), slevel, fmt.Sprintf(msg, args...))
}
return true
}
}
// Context is a default rule context used in teleport
type Context struct {
// User is currently authenticated user
User UserState
// Resource is an optional resource, in case if the rule
// checks access to the resource
Resource types.Resource
// Resource153 is an optional resource, in case if the rule
// checks access to the resource
Resource153 types.Resource153
// Session is an optional session.end or windows.desktop.session.end event.
// These events hold information about session recordings.
Session events.AuditEvent
// SSHSession is an optional (active) SSH session.
SSHSession *session.Session
// HostCert is an optional host certificate.
HostCert *HostCertContext
// SessionTracker is an optional session tracker, in case if the rule checks access to the tracker.
SessionTracker types.SessionTracker
// AccessChecker is an optional access checker that can be used
// to check access to other resources.
AccessChecker AccessChecker
}
// String returns user friendly representation of this context
func (ctx *Context) String() string {
return fmt.Sprintf("user %v, resource: %v, resource153: %v", ctx.User, ctx.Resource, ctx.Resource153)
}
const (
// UserIdentifier represents user registered identifier in the rules
UserIdentifier = "user"
// ResourceIdentifier represents resource registered identifier in the rules
ResourceIdentifier = "resource"
// ResourceLabelsIdentifier refers to the static and dynamic labels in a resource.
ResourceLabelsIdentifier = "labels"
// ResourceNameIdentifier refers to two different fields depending on the kind of resource:
// - KindNode will refer to its resource.spec.hostname field
// - All other kinds will refer to its resource.metadata.name field
// It refers to two different fields because the way this shorthand is being used,
// implies it will return the name of the resource where users identifies nodes
// by its hostname and all other resources that can be `ls` queried is identified
// by its metadata name.
ResourceNameIdentifier = "name"
// SessionIdentifier refers to a session (recording) in the rules.
SessionIdentifier = "session"
// SSHSessionIdentifier refers to an (active) SSH session in the rules.
SSHSessionIdentifier = "ssh_session"
// ImpersonateRoleIdentifier is a role to impersonate
ImpersonateRoleIdentifier = "impersonate_role"
// ImpersonateUserIdentifier is a user to impersonate
ImpersonateUserIdentifier = "impersonate_user"
// HostCertIdentifier refers to a host certificate being created.
HostCertIdentifier = "host_cert"
// SessionTrackerIdentifier refers to a session tracker in the rules.
SessionTrackerIdentifier = "session_tracker"
)
// GetResource returns resource specified in the context,
// returns error if not specified.
func (ctx *Context) GetResource() (types.Resource, error) {
switch {
case ctx.Resource == nil && ctx.Resource153 == nil:
return nil, trace.NotFound("resource is not set in the context")
case ctx.Resource == nil && ctx.Resource153 != nil:
return types.Resource153ToLegacy(ctx.Resource153), nil
case ctx.Resource != nil && ctx.Resource153 == nil:
return ctx.Resource, nil
default:
return nil, trace.BadParameter("only one resource should be provided")
}
}
func (ctx *Context) GetAccessChecker() (AccessChecker, error) {
if ctx.AccessChecker == nil {
return nil, trace.NotFound("access checker is not set in the context")
}
return ctx.AccessChecker, nil
}
// GetIdentifier returns identifier defined in a context
func (ctx *Context) GetIdentifier(fields []string) (any, error) {
switch fields[0] {
case UserIdentifier:
var user UserState
if ctx.User == nil {
user = emptyUser
} else {
user = ctx.User
}
return predicate.GetFieldByTag(user, teleport.JSON, fields[1:])
case ResourceIdentifier:
var r any
switch {
case ctx.Resource == nil && ctx.Resource153 == nil:
r = emptyResource
case ctx.Resource == nil && ctx.Resource153 != nil:
r = ctx.Resource153
case ctx.Resource != nil && ctx.Resource153 == nil:
r = ctx.Resource
default:
return nil, trace.BadParameter("only one resource should be provided")
}
return predicate.GetFieldByTag(r, teleport.JSON, fields[1:])
case SessionIdentifier:
var session events.AuditEvent = &events.SessionEnd{}
switch ctx.Session.(type) {
case *events.SessionEnd, *events.WindowsDesktopSessionEnd, *events.LinuxDesktopSessionEnd, *events.DatabaseSessionEnd, *events.AppSessionChunk:
session = ctx.Session
}
v, origErr := predicate.GetFieldByTag(session, teleport.JSON, fields[1:])
if trace.IsNotFound(origErr) {
// Special case: session is a special resource because
// it's backed by different object kinds (events.SessionEnd,
// events.WindowsDesktopSessionEnd, events.DatabaseSessionEnd)
// and these objects have different schemas, so it's possible that
// the parser can't find a field mentioned in the "where" clause.
// In this case, we try to find the field in all supported
// session end events and return the value if found.
if v, err := getMissingEmptyFieldForSessionEnd(fields); err == nil {
return v, nil
}
}
return v, trace.Wrap(origErr)
case SSHSessionIdentifier:
// Do not expose the original session.Session, instead transform it into a
// ctxSession so the exposed fields match our desired API.
return predicate.GetFieldByTag(toCtxSession(ctx.SSHSession), teleport.JSON, fields[1:])
case HostCertIdentifier:
var hostCert *HostCertContext
if ctx.HostCert == nil {
hostCert = emptyHostCert
} else {
hostCert = ctx.HostCert
}
return predicate.GetFieldByTag(hostCert, teleport.JSON, fields[1:])
case SessionTrackerIdentifier:
return predicate.GetFieldByTag(toCtxTracker(ctx.SessionTracker), teleport.JSON, fields[1:])
case ResourceNameIdentifier:
if len(fields) > 1 {
return nil, trace.BadParameter(
"only one field is supported with identifier %q, got %d: %v",
ResourceNameIdentifier,
len(fields),
fields,
)
}
switch {
case ctx.Resource == nil && ctx.Resource153 == nil:
return "", nil
case ctx.Resource == nil && ctx.Resource153 != nil:
return ctx.Resource153.GetMetadata().GetName(), nil
case ctx.Resource != nil && ctx.Resource153 == nil:
return ctx.Resource.GetName(), nil
default:
return nil, trace.BadParameter("only one resource should be provided")
}
default:
return nil, trace.NotFound("%v is not defined", strings.Join(fields, "."))
}
}
func getMissingEmptyFieldForSessionEnd(fields []string) (any, error) {
for _, emptySession := range []events.AuditEvent{&events.SessionEnd{}, &events.WindowsDesktopSessionEnd{}, &events.DatabaseSessionEnd{}, &events.AppSessionChunk{}} {
v, err := predicate.GetFieldByTag(emptySession, teleport.JSON, fields[1:])
if err == nil {
return v, nil
}
}
return nil, trace.NotFound("field %q is not found in any supported session end event", strings.Join(fields, "."))
}
// ctxSession represents the public contract of a session.Session, as exposed
// to a Context rule.
// See RFD 82: https://github.com/gravitational/teleport/blob/master/rfd/0082-session-tracker-resource-rbac.md
type ctxTracker struct {
SessionID string `json:"session_id"`
Kind string `json:"kind"`
Participants []string `json:"participants"`
State string `json:"state"`
Hostname string `json:"hostname"`
Address string `json:"address"`
Login string `json:"login"`
Cluster string `json:"cluster"`
KubeCluster string `json:"kube_cluster"`
HostUser string `json:"host_user"`
HostRoles []string `json:"host_roles"`
}
func toCtxTracker(t types.SessionTracker) ctxTracker {
if t == nil {
return ctxTracker{}
}
getParticipants := func(s types.SessionTracker) []string {
participants := s.GetParticipants()
names := make([]string, len(participants))
for i, participant := range participants {
// Participant for RBAC must be represented as `remote-{user}-{cluster}`.
// if they belong to a different cluster. This is because the user
// is also named like that when they authenticate.
names[i] = UsernameForCluster(UsernameForClusterConfig{
User: participant.User,
OriginClusterName: participant.Cluster,
LocalClusterName: s.GetClusterName(),
})
}
return names
}
getHostRoles := func(s types.SessionTracker) []string {
policySets := s.GetHostPolicySets()
roles := make([]string, len(policySets))
for i, policySet := range policySets {
roles[i] = policySet.Name
}
return roles
}
return ctxTracker{
SessionID: t.GetSessionID(),
Kind: t.GetKind(),
Participants: getParticipants(t),
State: string(t.GetState()),
Hostname: t.GetHostname(),
Address: t.GetAddress(),
Login: t.GetLogin(),
Cluster: t.GetClusterName(),
KubeCluster: t.GetKubeCluster(),
HostUser: t.GetHostUser(),
HostRoles: getHostRoles(t),
}
}
// ctxSession represents the public contract of a session.Session, as exposed
// to a Context rule.
// See RFD 45:
// https://github.com/gravitational/teleport/blob/master/rfd/0045-ssh_session-where-condition.md#replacing-parties-by-usernames.
type ctxSession struct {
// Namespace is a session namespace, separating sessions from each other.
Namespace string `json:"namespace"`
// Login is a login used by all parties joining the session.
Login string `json:"login"`
// Created records the information about the time when session was created.
Created time.Time `json:"created"`
// LastActive holds the information about when the session was last active.
LastActive time.Time `json:"last_active"`
// ServerID of session.
ServerID string `json:"server_id"`
// ServerHostname of session.
ServerHostname string `json:"server_hostname"`
// ServerAddr of session.
ServerAddr string `json:"server_addr"`
// ClusterName is the name of cluster that this session belongs to.
ClusterName string `json:"cluster_name"`
// Participants is a list of session participants expressed as usernames.
Participants []string `json:"participants"`
}
func toCtxSession(s *session.Session) ctxSession {
if s == nil {
return ctxSession{}
}
return ctxSession{
Namespace: s.Namespace,
Login: s.Login,
Created: s.Created,
LastActive: s.LastActive,
ServerID: s.ServerID,
ServerHostname: s.ServerHostname,
ServerAddr: s.ServerAddr,
ClusterName: s.ClusterName,
Participants: s.Participants(),
}
}
// HostCertContext is used to evaluate the `where` condition on a `host_cert`
// pseudo-resource. These resources only exist for RBAC purposes and do not
// exist in the database.
type HostCertContext struct {
// HostID is the host ID in the cert request.
HostID string `json:"host_id"`
// NodeName is the node name in the cert request.
NodeName string `json:"node_name"`
// Principals is the list of requested certificate principals.
Principals []string `json:"principals"`
// ClusterName is the name of the cluster for which the certificate should
// be issued.
ClusterName string `json:"cluster_name"`
// Role is the name of the Teleport role for which the cert should be
// issued.
Role types.SystemRole `json:"role"`
// TTL is the requested certificate TTL.
TTL time.Duration `json:"ttl"`
}
// emptyResource is used when no resource is specified
var emptyResource = &EmptyResource{}
// emptyUser is used when no user is specified
var emptyUser = &types.UserV2{}
// emptyHostCert is an empty host certificate used when no host cert is
// specified
var emptyHostCert = &HostCertContext{}
// EmptyResource is used to represent a use case when no resource
// is specified in the rules matcher
type EmptyResource struct {
// Kind is a resource kind
Kind string `json:"kind"`
// SubKind is a resource sub kind
SubKind string `json:"sub_kind,omitempty"`
// Version is a resource version
Version string `json:"version"`
// Metadata is Role metadata
Metadata types.Metadata `json:"metadata"`
}
// GetVersion returns resource version
func (r *EmptyResource) GetVersion() string {
return r.Version
}
// GetSubKind returns resource sub kind
func (r *EmptyResource) GetSubKind() string {
return r.SubKind
}
// SetSubKind sets resource subkind
func (r *EmptyResource) SetSubKind(s string) {
r.SubKind = s
}
// GetKind returns resource kind
func (r *EmptyResource) GetKind() string {
return r.Kind
}
// GetRevision returns the revision
func (r *EmptyResource) GetRevision() string {
return r.Metadata.GetRevision()
}
// SetRevision sets the revision
func (r *EmptyResource) SetRevision(rev string) {
r.Metadata.SetRevision(rev)
}
// SetExpiry sets expiry time for the object.
func (r *EmptyResource) SetExpiry(expires time.Time) {
r.Metadata.SetExpiry(expires)
}
// Expiry returns the expiry time for the object.
func (r *EmptyResource) Expiry() time.Time {
return r.Metadata.Expiry()
}
// SetName sets the role name and is a shortcut for SetMetadata().Name.
func (r *EmptyResource) SetName(s string) {
r.Metadata.Name = s
}
// GetName gets the role name and is a shortcut for GetMetadata().Name.
func (r *EmptyResource) GetName() string {
return r.Metadata.Name
}
// GetMetadata returns role metadata.
func (r *EmptyResource) GetMetadata() types.Metadata {
return r.Metadata
}
func (r *EmptyResource) CheckAndSetDefaults() error { return nil }
// newParserForIdentifierSubcondition returns a parser customized for
// extracting the largest admissible subexpression of a `where` condition that
// involves the given identifier.
//
// For example, consider the `where` condition
// `contains(session.participants, "user") && equals(user.metadata.name, "user")`.
// Given a RuleContext where user.metadata.name is equal to "user", its largest
// admissible subcondition involving the identifier "session" is
// `contains(session.participants, "user")`. With another RuleContext the
// largest such subcondition is the empty expression.
func newParserForIdentifierSubcondition(ctx RuleContext, identifier string) (predicate.Parser, error) {
binaryPred := func(predFn func(a, b any) predicate.BoolPredicate, exprFn func(a, b types.WhereExpr) types.WhereExpr) func(a, b any) types.WhereExpr {
return func(a, b any) types.WhereExpr {
an, aOK := a.(types.WhereExpr)
if !aOK {
an = types.WhereExpr{Literal: a}
}
bn, bOK := b.(types.WhereExpr)
if !bOK {
bn = types.WhereExpr{Literal: b}
}
if an.Literal != nil && bn.Literal != nil {
return types.WhereExpr{Literal: predFn(an.Literal, bn.Literal)()}
}
return exprFn(an, bn)
}
}
return predicate.NewParser(predicate.Def{
Operators: predicate.Operators{
AND: func(a, b types.WhereExpr) types.WhereExpr {
aVal, aOK := a.Literal.(bool)
bVal, bOK := b.Literal.(bool)
switch {
case aOK && bOK:
return types.WhereExpr{Literal: aVal && bVal}
case aVal:
return b
case bVal:
return a
case aOK || bOK:
return types.WhereExpr{Literal: false}
default:
return types.WhereExpr{And: types.WhereExpr2{L: &a, R: &b}}
}
},
OR: func(a, b types.WhereExpr) types.WhereExpr {
aVal, aOK := a.Literal.(bool)
bVal, bOK := b.Literal.(bool)
switch {
case aOK && bOK:
return types.WhereExpr{Literal: aVal || bVal}
case aVal || bVal:
return types.WhereExpr{Literal: true}
case aOK:
return b
case bOK:
return a
default:
return types.WhereExpr{Or: types.WhereExpr2{L: &a, R: &b}}
}
},
NOT: func(expr types.WhereExpr) types.WhereExpr {
if val, ok := expr.Literal.(bool); ok {
return types.WhereExpr{Literal: !val}
}
return types.WhereExpr{Not: &expr}
},
EQ: func(a, b any) types.WhereExpr {
aExpr, ok := a.(types.WhereExpr)
if !ok {
aExpr = types.WhereExpr{Literal: a}
}
bExpr, ok := b.(types.WhereExpr)
if !ok {
bExpr = types.WhereExpr{Literal: b}
}
return types.WhereExpr{Equals: types.WhereExpr2{L: &aExpr, R: &bExpr}}
},
NEQ: func(a, b any) types.WhereExpr {
aExpr, ok := a.(types.WhereExpr)
if !ok {
aExpr = types.WhereExpr{Literal: a}
}
bExpr, ok := b.(types.WhereExpr)
if !ok {
bExpr = types.WhereExpr{Literal: b}
}
return types.WhereExpr{
Not: &types.WhereExpr{Equals: types.WhereExpr2{L: &aExpr, R: &bExpr}},
}
},
},
Functions: map[string]any{
"equals": binaryPred(predicate.Equals, func(a, b types.WhereExpr) types.WhereExpr {
return types.WhereExpr{Equals: types.WhereExpr2{L: &a, R: &b}}
}),
"contains": binaryPred(predicate.Contains, func(a, b types.WhereExpr) types.WhereExpr {
return types.WhereExpr{Contains: types.WhereExpr2{L: &a, R: &b}}
}),
"contains_all": binaryPred(predicateContainsAll, func(a, b types.WhereExpr) types.WhereExpr {
return types.WhereExpr{ContainsAll: types.WhereExpr2{L: &a, R: &b}}
}),
"contains_any": binaryPred(predicateContainsAny, func(a, b types.WhereExpr) types.WhereExpr {
return types.WhereExpr{ContainsAny: types.WhereExpr2{L: &a, R: &b}}
}),
"set": func(a ...any) types.WhereExpr {
aVal := make([]string, 0, len(a))
for _, v := range a {
if str, ok := v.(string); ok {
aVal = append(aVal, str)
}
}
return types.WhereExpr{Literal: aVal}
},
"can_view": func() types.WhereExpr {
return types.WhereExpr{CanView: &types.WhereNoExpr{}}
},
},
GetIdentifier: func(fields []string) (any, error) {
if fields[0] == identifier {
// TODO: Session events have only one level of attributes. Support for
// more nested levels may be added when needed for other objects.
if len(fields) != 2 {
return nil, trace.BadParameter("only exactly two fields are supported with identifier %q, got %d: %v", identifier, len(fields), fields)
}
return types.WhereExpr{Field: fields[1]}, nil
}
lit, err := ctx.GetIdentifier(fields)
return types.WhereExpr{Literal: lit}, trace.Wrap(err)
},
GetProperty: func(mapVal, keyVal any) (any, error) {
mapExpr, mapOK := mapVal.(types.WhereExpr)
if !mapOK {
mapExpr = types.WhereExpr{Literal: mapVal}
}
keyExpr, keyOK := keyVal.(types.WhereExpr)
if !keyOK {
keyExpr = types.WhereExpr{Literal: keyVal}
}
if mapExpr.Field != "" && keyExpr.Literal != nil {
return types.WhereExpr{
MapRef: &types.WhereExpr2{
L: &mapExpr,
R: &keyExpr,
},
}, nil
}
if mapExpr.Literal == nil || keyExpr.Literal == nil {
// TODO: Add support for general WhereExpr.
return nil, trace.BadParameter("GetProperty is implemented only for literals")
}
return GetStringMapValue(mapExpr.Literal, keyExpr.Literal)
},
})
}
// predicateContainsAll is a custom function to test if all entries in a []string
// are contained in another []string. Order does not matter, but all entries
// in the second slice must be present in the first slice.
func predicateContainsAll(a, b any) predicate.BoolPredicate {
return func() bool {
aval, ok := a.([]string)
if !ok {
return false
}
bval, ok := b.([]string)
if !ok {
return false
}
if len(aval) == 0 || len(bval) == 0 {
return false
}
aSet := set.New(aval...)
for _, v := range bval {
if !aSet.Contains(v) {
return false
}
}
return true
}
}
// predicateContainsAny is a custom function to test if any entry in a []string
func predicateContainsAny(a, b any) predicate.BoolPredicate {
return func() bool {
aval, ok := a.([]string)
if !ok {
return false
}
bval, ok := b.([]string)
if !ok {
return false
}
if len(aval) == 0 || len(bval) == 0 {
return false
}
aSet := set.New(aval...)
for _, v := range bval {
if aSet.Contains(v) {
return true
}
}
return false
}
}
type lazyStringSplit struct {
value, delimiter string
}
// Slice implements [typical.StringSlice].
func (s *lazyStringSplit) Slice() []string {
return strings.Split(s.value, s.delimiter)
}
// Contains implements [typical.StringSlice].
func (s *lazyStringSplit) Contains(target string) bool {
return splitContains(s.value, s.delimiter, target)
}
func splitContains(value, delimiter, target string) bool {
if delimiter == "" {
// an empty delimiter was unfortunately not preemptively disallowed, so
// we have to match the behavior of [strings.Split] with an empty
// delimiter, i.e. splitting up the string in individual utf8 sequences
// or single bytes of non-well-formed utf8
if target == "" {
// exploding a string into utf8 sequences never yields an empty
// string, so we can't ever match an empty target if the delimiter
// is empty
return false
}
if value == "" {
// exploding the empty string into utf8 sequences yields nothing
return false
}
r, s := utf8.DecodeRuneInString(target)
if len(target) != s {
// target is more than one rune or invalid non-utf8 byte, so it will
// never match anything yielded by SplitSeq(value, "")
return false
}
if s != 1 || r != utf8.RuneError {
// target is a well formed utf8 sequence, so we can check if it's in
// value by means of Contains
return strings.Contains(value, target)
}
// target is a single byte that's not a valid utf8 sequence but it might
// appear in value as part of a valid utf8 sequence, so we can't use a
// bytewise contains, unfortunately, and the predicate language supports
// go escaping for strings, so we can't rely on the well-formedness of
// strings either; if we hit this case we just fall back to the slower
// case using SplitSeq
} else if value == "" {
// the empty value splits into a single empty string since the delimiter
// is nonempty
return target == ""
} else if strings.Contains(target, delimiter) || !strings.Contains(value, target) {
// delimiter is nonempty, so SplitSeq works bytewise, and there's no way
// we'll ever get a match if the target contains the delimiter or if the
// target is not contained in the value; note that this check is
// wasteful in case of a successful match, but erring on the side of
// excluding matches will hopefully result in overall faster behavior,
// especially when resource expressions are used as a predicate to
// select one or few resources
return false
}
for v := range strings.SplitSeq(value, delimiter) {
if target == v {
return true
}
}
return false
}
func splitContainsSinglebyteAffix(value, delimTargetDelim string) bool {
// the delimiter is one byte and it's at the head and tail of
// delimTargetDelim, so if we don't have two bytes we can't do some of the
// slicing; this is an internal-use function so we're not really concerned
// with sanity checks - if called incorrectly, this function will return a
// useless value but it will not crash
if len(delimTargetDelim) < 2 {
return false
}
if target := delimTargetDelim[1 : len(delimTargetDelim)-1]; value == target {
// ("foo", ",foo,")
return true
}
if targetDelim := delimTargetDelim[1:]; strings.HasPrefix(value, targetDelim) {
// ("foo,bar", ",foo,")
return true
}
if delimTarget := delimTargetDelim[:len(delimTargetDelim)-1]; strings.HasSuffix(value, delimTarget) {
// ("bar,foo", ",foo,")
return true
}
// ("bar,qux,foo,baz", ",foo,")
return strings.Contains(value, delimTargetDelim)
}
func newResourceExpressionParser(opts ...func(*typical.ParserSpec[types.ResourceWithLabels])) (*typical.Parser[types.ResourceWithLabels, bool], error) {
spec := typical.ParserSpec[types.ResourceWithLabels]{
Variables: map[string]typical.Variable{
"resource.metadata.labels": typical.DynamicVariable(func(r types.ResourceWithLabels) (map[string]string, error) {
return r.GetStaticLabels(), nil
}),
"resource.metadata.name": typical.DynamicVariable(func(r types.ResourceWithLabels) (string, error) {
return r.GetName(), nil
}),
"labels": typical.DynamicMapFunction(func(r types.ResourceWithLabels, key string) (string, error) {
val, _ := r.GetLabel(key)
return val, nil
}),
"name": typical.DynamicVariable(func(r types.ResourceWithLabels) (string, error) {
// For nodes, the resource "name" that user expects is the
// nodes hostname, not its UUID. Currently, for other resources,
// the metadata.name returns the name as expected.
if server, ok := r.(types.Server); ok {
return server.GetHostname(), nil
}
return r.GetName(), nil
}),
"health.status": typical.DynamicVariable(func(r types.ResourceWithLabels) (string, error) {
if r, ok := r.(types.TargetHealthStatusGetter); ok {
return string(r.GetTargetHealthStatus()), nil
}
return "", nil
}),
},
Functions: map[string]typical.Function{
"hasPrefix": typical.BinaryFunction[types.ResourceWithLabels](func(s, suffix string) (bool, error) {
return strings.HasPrefix(s, suffix), nil
}),
"equals": typical.BinaryFunction[types.ResourceWithLabels](func(a, b string) (bool, error) {
return strings.Compare(a, b) == 0, nil
}),
"search": typical.UnaryVariadicFunctionWithEnv(func(r types.ResourceWithLabels, v ...string) (bool, error) {
return r.MatchSearch(v), nil
}),
"exists": typical.UnaryFunction[types.ResourceWithLabels](func(value string) (bool, error) {
return value != "", nil
}),
"split": typical.BinaryFunction[types.ResourceWithLabels](func(value string, delimiter string) (typical.StringSlice, error) {
return &lazyStringSplit{value: value, delimiter: delimiter}, nil
}),
"contains": typical.BinaryFunction[types.ResourceWithLabels](func(list typical.StringSlice, value string) (bool, error) {
return list.Contains(value), nil
}),
"__split_contains": typical.TernaryFunction[types.ResourceWithLabels](func(value string, delimiter string, target string) (bool, error) {
return splitContains(value, delimiter, target), nil
}),
"__split_contains_singlebyte_affix": typical.BinaryFunction[types.ResourceWithLabels](func(value string, delimTargetDelim string) (bool, error) {
return splitContainsSinglebyteAffix(value, delimTargetDelim), nil
}),
},
GetUnknownIdentifier: func(env types.ResourceWithLabels, fields []string) (any, error) {
if fields[0] == ResourceIdentifier {
if f, err := predicate.GetFieldByTag(env, teleport.JSON, fields[1:]); err == nil {
return f, nil
}
}
identifier := strings.Join(fields, ".")
return nil, trace.BadParameter("identifier %s is not defined", identifier)
},
}
for _, opt := range opts {
opt(&spec)
}
parser, err := typical.NewParser[types.ResourceWithLabels, bool](spec)
if err != nil {
return nil, trace.Wrap(err)
}
return parser, nil
}
var resourceExpressionParserOnce sync.Once
var resourceExpressionParser *typical.Parser[types.ResourceWithLabels, bool]
func getResourceExpressionParser() (*typical.Parser[types.ResourceWithLabels, bool], error) {
resourceExpressionParserOnce.Do(func() {
p, err := newResourceExpressionParser()
if err != nil {
return
}
resourceExpressionParser = p
})
if resourceExpressionParser != nil {
return resourceExpressionParser, nil
}
// if there was an error creating the parser, which is impossible at the
// time of writing, just try again on every call so we can forward the error
// to the caller
return newResourceExpressionParser()
}
// NewResourceExpression returns a [typical.Expression] that is to be evaluated against a
// [types.ResourceWithLabels]. It is customized to allow short identifiers common in all
// resources:
// - shorthand `name` refers to `resource.spec.hostname` for node resources, or it refers
// to `resource.metadata.name` for all other resources eg: `name == "app-name-jenkins"`
// - shorthand `labels` refers to resource `resource.metadata.labels + resource.spec.dynamic_labels`
// eg: `labels.env == "prod"`
//
// All other fields can be referenced by starting expression with identifier `resource`
// followed by the names of the json fields ie: `resource.spec.public_addr`.
func NewResourceExpression(expression string) (typical.Expression[types.ResourceWithLabels, bool], error) {
parser, err := getResourceExpressionParser()
if err != nil {
return nil, trace.Wrap(err)
}
return newResourceExpression(expression, parser)
}
func newResourceExpression(expression string, parser *typical.Parser[types.ResourceWithLabels, bool]) (typical.Expression[types.ResourceWithLabels, bool], error) {
astExpr, err := goparser.ParseExpr(expression)
if err != nil {
return nil, trace.Wrap(err)
}
for _, postFunc := range []astutil.ApplyFunc{
combineSplitContains,
optimizeSplitContains,
} {
newNode := astutil.Apply(astExpr, nil, postFunc)
astExpr, _ = newNode.(ast.Expr)
if astExpr == nil {
return nil, trace.Errorf("expected ast.Expr from AST optimization step, got %T (this is a bug)", newNode)
}
}
expr, err := parser.ParseAST(astExpr)
return expr, trace.Wrap(err)
}
// combineSplitContains is an [astutil.ApplyFunc] that replaces the combination
// of calls of "contains(split(value, delimiter), target)" into a single call to
// "__split_contains(value, delimiter, target)".
func combineSplitContains(cursor *astutil.Cursor) (cont bool) {
// we don't terminate the AST visit early so we always return true
cont = true
containsCall := getIdentCall(cursor.Node(), "contains")
if containsCall == nil {
return
}
if len(containsCall.Args) != 2 {
return
}
splitCall := getIdentCall(ast.Unparen(containsCall.Args[0]), "split")
if splitCall == nil {
return
}
if len(splitCall.Args) != 2 {
return
}
// contains(split(value, delim), target) -> __split_contains(value, delim, target)
cursor.Replace(&ast.CallExpr{
Fun: &ast.Ident{Name: "__split_contains"},
Args: []ast.Expr{
ast.Unparen(splitCall.Args[0]),
ast.Unparen(splitCall.Args[1]),
ast.Unparen(containsCall.Args[1]),
},
})
return
}
func isImpossibleSplitContainsMatch(target, delimiter string) bool {
return delimiter != "" && strings.Contains(target, delimiter)
}
// optimizeSplitContains is an [astutil.ApplyFunc] that replaces calls to
// __split_contains where the delimiter is a nonempty literal string and the
// target is a literal string into calls to __split_contains_singlebyte_affix,
// by combining the delimiter and literal strings into a new literal string that
// can be more efficiently matched.
func optimizeSplitContains(cursor *astutil.Cursor) (cont bool) {
// we don't terminate the AST visit early so we always return true
cont = true
splitContainsCall := getIdentCall(cursor.Node(), "__split_contains")
if splitContainsCall == nil {
return
}
if len(splitContainsCall.Args) != 3 {
return
}
delimiter := getLiteralString(ast.Unparen(splitContainsCall.Args[1]))
target := getLiteralString(ast.Unparen(splitContainsCall.Args[2]))
if delimiter == nil || target == nil {
return
}
if isImpossibleSplitContainsMatch(*target, *delimiter) {
// impossible match, turn it into a function call that unconditionally
// returns false very quickly
// TODO(espadolini): figure out if adding a "false" or "__false"
// variable has the potential to break things
// __split_contains(value, ",", "contains,delimiter") -> __split_contains_singlebyte_affix(value, "")
cursor.Replace(&ast.CallExpr{
Fun: &ast.Ident{Name: "__split_contains_singlebyte_affix"},
Args: []ast.Expr{
ast.Unparen(splitContainsCall.Args[0]),
&ast.BasicLit{Kind: token.STRING, Value: `""`},
},
})
return
}
if len(*delimiter) != 1 {
return
}
// the delimiter is a one byte literal string and the target is a literal
// string, so we affix the delimiter to the head and tail of the target
// string so we can just call [strings.Contains] which is almost surely much
// more efficient than anything else we could do; the delimiter has to be
// single byte because __split_contains_singlebyte_affix must be able to
// skip it when checking for the target being at the very beginning or the
// very end of the string, and we don't want to copy every string in the
// slice just to add the delimiter at the beginning and end for that
// __split_contains(sliceExpr, ",", "literal") -> __split_contains_singlebyte_affix(sliceExpr, ",literal,")
cursor.Replace(&ast.CallExpr{
Fun: &ast.Ident{Name: "__split_contains_singlebyte_affix"},
Args: []ast.Expr{
ast.Unparen(splitContainsCall.Args[0]),
&ast.BasicLit{
Kind: token.STRING,
Value: strconv.Quote(*delimiter + *target + *delimiter),
},
},
})
return
}
func getIdentCall(n ast.Node, ident string) *ast.CallExpr {
c, _ := n.(*ast.CallExpr)
if c == nil {
return nil
}
// predicate doesn't handle function expressions, not even to unparen them,
// so `(foo)(1, 2, 3)` is not valid and we only need to check for
// [*ast.Ident]
if i, _ := c.Fun.(*ast.Ident); i == nil || i.Name != ident {
return nil
}
return c
}
func getLiteralString(n ast.Node) *string {
l, _ := n.(*ast.BasicLit)
if l == nil {
return nil
}
if l.Kind != token.STRING {
return nil
}
s, err := strconv.Unquote(l.Value)
if err != nil {
return nil
}
return &s
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// PluginDataGetter defines the interface for getting plugin data.
type PluginDataGetter interface {
// GetPluginData loads all plugin data matching the supplied filter.
GetPluginData(ctx context.Context, filter types.PluginDataFilter) ([]types.PluginData, error)
}
// PluginData defines the interface for managing plugin data.
type PluginData interface {
PluginDataGetter
// UpdatePluginData updates a per-resource PluginData entry.
UpdatePluginData(ctx context.Context, params types.PluginDataUpdateParams) error
}
// MarshalPluginData marshals the PluginData resource to JSON.
func MarshalPluginData(pluginData types.PluginData, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch pluginData := pluginData.(type) {
case *types.PluginDataV3:
if err := pluginData.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, pluginData))
default:
return nil, trace.BadParameter("unrecognized plugin data type: %T", pluginData)
}
}
// UnmarshalPluginData unmarshals the PluginData resource from JSON.
func UnmarshalPluginData(raw []byte, opts ...MarshalOption) (types.PluginData, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var data types.PluginDataV3
if err := utils.FastUnmarshal(raw, &data); err != nil {
return nil, trace.Wrap(err)
}
if err := data.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
data.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
data.SetExpiry(cfg.Expires)
}
return &data, nil
}
/*
* 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/trace"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/protoadapt"
"github.com/gravitational/teleport/api/types"
)
// PluginStaticCredentials is the plugin static credentials service
type PluginStaticCredentials interface {
// CreatePluginStaticCredentials will create a new plugin static credentials resource.
CreatePluginStaticCredentials(ctx context.Context, pluginStaticCredentials types.PluginStaticCredentials) error
// GetPluginStaticCredentials will get a plugin static credentials resource by name.
GetPluginStaticCredentials(ctx context.Context, name string) (types.PluginStaticCredentials, error)
// GetPluginStaticCredentialsByLabels will get a list of plugin static credentials resource by matching labels.
GetPluginStaticCredentialsByLabels(ctx context.Context, labels map[string]string) ([]types.PluginStaticCredentials, error)
// UpdatePluginStaticCredentials will update a plugin static credentials' resource.
UpdatePluginStaticCredentials(ctx context.Context, pluginStaticCredentials types.PluginStaticCredentials) (types.PluginStaticCredentials, error)
// DeletePluginStaticCredentials will delete a plugin static credentials resource.
DeletePluginStaticCredentials(ctx context.Context, name string) error
// GetAllPluginStaticCredentials will get all plugin static credentials.
GetAllPluginStaticCredentials(ctx context.Context) ([]types.PluginStaticCredentials, error)
}
// MarshalPluginStaticCredentials marshals PluginStaticCredentials resource to JSON.
func MarshalPluginStaticCredentials(pluginStaticCredentials types.PluginStaticCredentials, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch pluginStaticCredentials := pluginStaticCredentials.(type) {
case *types.PluginStaticCredentialsV1:
if err := pluginStaticCredentials.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
data, err := protojson.Marshal(protoadapt.MessageV2Of(maybeResetProtoRevision(cfg.PreserveRevision, pluginStaticCredentials)))
if err != nil {
return nil, trace.Wrap(err)
}
return data, nil
default:
return nil, trace.BadParameter("unsupported plugin static credentials resource %T", pluginStaticCredentials)
}
}
// UnmarshalPluginStaticCredentials unmarshals the plugin static credentials resource from JSON.
func UnmarshalPluginStaticCredentials(data []byte, opts ...MarshalOption) (types.PluginStaticCredentials, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing plugin static credentials resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.MessageWithHeader
// every field but one is unknown to [types.MessageWithHeader] so this
// unmarshal must discard unknown fields
if err := (protojson.UnmarshalOptions{DiscardUnknown: true}).Unmarshal(data, protoadapt.MessageV2Of(&h)); err != nil {
return nil, trace.BadParameter("%s", err)
}
switch h.ResourceHeader.Version {
case types.V1:
var pluginStaticCredentials types.PluginStaticCredentialsV1
if err := (protojson.UnmarshalOptions{DiscardUnknown: !cfg.DisallowUnknown}).Unmarshal(data, protoadapt.MessageV2Of(&pluginStaticCredentials)); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := pluginStaticCredentials.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
pluginStaticCredentials.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
pluginStaticCredentials.SetExpiry(cfg.Expires)
}
return &pluginStaticCredentials, nil
}
return nil, trace.BadParameter("unsupported plugin static credentials resource version %q", h.Version)
}
/*
* 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 (
"bytes"
"context"
"github.com/gogo/protobuf/jsonpb" //nolint:depguard // needed for backwards compatibility
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
type PluginGetter interface {
GetPlugin(ctx context.Context, name string, withSecrets bool) (types.Plugin, error)
GetPlugins(ctx context.Context, withSecrets bool) ([]types.Plugin, error)
ListPlugins(ctx context.Context, limit int, startKey string, withSecrets bool) ([]types.Plugin, string, error)
HasPluginType(ctx context.Context, pluginType types.PluginType) (bool, error)
}
// Plugins is the plugin service
type Plugins interface {
PluginGetter
CreatePlugin(ctx context.Context, plugin types.Plugin) error
UpdatePlugin(ctx context.Context, plugin types.Plugin) (types.Plugin, error)
DeleteAllPlugins(ctx context.Context) error
DeletePlugin(ctx context.Context, name string) error
SetPluginCredentials(ctx context.Context, name string, creds types.PluginCredentials) error
SetPluginStatus(ctx context.Context, name string, creds types.PluginStatus) error
}
// MarshalPlugin marshals Plugin resource to JSON.
func MarshalPlugin(plugin types.Plugin, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch plugin := plugin.(type) {
case *types.PluginV1:
var buf bytes.Buffer
err := (&jsonpb.Marshaler{}).Marshal(&buf, maybeResetProtoRevision(cfg.PreserveRevision, plugin))
if err != nil {
return nil, trace.Wrap(err)
}
return buf.Bytes(), nil
default:
return nil, trace.BadParameter("unsupported plugin resource %T", plugin)
}
}
// UnmarshalPlugin unmarshals the plugin resource from JSON.
func UnmarshalPlugin(data []byte, opts ...MarshalOption) (types.Plugin, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing plugin resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var plugin types.PluginV1
m := jsonpb.Unmarshaler{AllowUnknownFields: true}
if err := m.Unmarshal(bytes.NewReader(data), &plugin); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := plugin.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
plugin.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
plugin.SetExpiry(cfg.Expires)
}
return &plugin, nil
}
return nil, trace.BadParameter("unsupported plugin resource version %q", h.Version)
}
/*
* 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"
"fmt"
"log/slog"
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
healthcheckconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/healthcheckconfig/v1"
labelv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/label/v1"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/scopes/access"
"github.com/gravitational/teleport/lib/utils/set"
)
// NewSystemAutomaticAccessApproverRole creates a new Role that is allowed to
// approve any Access Request. This is restricted to Teleport Enterprise, and
// returns nil in non-Enterproise builds.
func NewSystemAutomaticAccessApproverRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.SystemAutomaticAccessApprovalRoleName,
Namespace: apidefaults.Namespace,
Description: "Approves any access request",
Labels: map[string]string{
types.TeleportInternalResourceType: types.SystemResource,
types.TeleportResourceRevision: "1",
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
ReviewRequests: &types.AccessReviewConditions{
Roles: []string{"*"},
},
},
},
}
role.CheckAndSetDefaults()
return role
}
// NewSystemAutomaticAccessBotUser returns a new User that has (via the
// the `PresetAutomaticAccessApprovalRoleName` role) the right to automatically
// approve any access requests.
//
// This user must not:
// - Be allowed to log into the cluster
// - Show up in user lists in WebUI
//
// TODO(tcsc): Implement/enforce above restrictions on this user
func NewSystemAutomaticAccessBotUser(buildType string) types.User {
if buildType != modules.BuildEnterprise {
return nil
}
user := &types.UserV2{
Kind: types.KindUser,
Version: types.V2,
Metadata: types.Metadata{
Name: teleport.SystemAccessApproverUserName,
Namespace: apidefaults.Namespace,
Description: "Used internally by Teleport to automatically approve access requests",
Labels: map[string]string{
types.TeleportInternalResourceType: string(types.SystemResource),
types.TeleportResourceRevision: "1",
},
},
Spec: types.UserSpecV2{
Roles: []string{teleport.SystemAutomaticAccessApprovalRoleName},
},
}
user.CheckAndSetDefaults()
return user
}
// NewPresetEditorRole returns a new pre-defined role for cluster
// editors who can edit cluster configuration resources.
func NewPresetEditorRole() types.Role {
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetEditorRoleName,
Namespace: apidefaults.Namespace,
Description: "Edit cluster configuration",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
Options: types.RoleOptions{
CertificateFormat: constants.CertificateFormatStandard,
MaxSessionTTL: types.NewDuration(apidefaults.MaxCertDuration),
SSHPortForwarding: &types.SSHPortForwarding{
Remote: &types.SSHRemotePortForwarding{
Enabled: types.NewBoolOption(true),
},
Local: &types.SSHLocalPortForwarding{
Enabled: types.NewBoolOption(true),
},
},
ForwardAgent: types.NewBool(true),
BPF: apidefaults.EnhancedEvents(),
RecordSession: &types.RecordSession{
Desktop: types.NewBoolOption(false),
},
},
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
Allow: types.RoleConditions{
Namespaces: []string{apidefaults.Namespace},
Rules: []types.Rule{
types.NewRule(types.KindUser, RW()),
types.NewRule(types.KindRole, RW()),
types.NewRule(types.KindBot, RW()),
types.NewRule(types.KindCrownJewel, RW()),
types.NewRule(types.KindDatabaseObjectImportRule, RW()),
types.NewRule(types.KindOIDC, RW()),
types.NewRule(types.KindSAML, RW()),
types.NewRule(types.KindGithub, RW()),
types.NewRule(types.KindOIDCRequest, RW()),
types.NewRule(types.KindSAMLRequest, RW()),
types.NewRule(types.KindGithubRequest, RW()),
types.NewRule(types.KindClusterAuditConfig, RW()),
types.NewRule(types.KindClusterAuthPreference, RW()),
types.NewRule(types.KindAuthConnector, RW()),
types.NewRule(types.KindClusterName, RW()),
types.NewRule(types.KindClusterNetworkingConfig, RW()),
types.NewRule(types.KindSessionRecordingConfig, RW()),
types.NewRule(types.KindExternalAuditStorage, RW()),
types.NewRule(types.KindUIConfig, RW()),
types.NewRule(types.KindTrustedCluster, RW()),
types.NewRule(types.KindRemoteCluster, RW()),
types.NewRule(types.KindToken, RW()),
types.NewRule(types.KindConnectionDiagnostic, RW()),
types.NewRule(types.KindDatabase, RW()),
types.NewRule(types.KindDatabaseCertificate, RW()),
types.NewRule(types.KindInstaller, RW()),
types.NewRule(types.KindDevice, append(RW(), types.VerbCreateEnrollToken, types.VerbEnroll)),
types.NewRule(types.KindMobileDevice, []string{types.VerbCreateEnrollToken}),
types.NewRule(types.KindDatabaseService, RO()),
types.NewRule(types.KindInstance, RO()),
types.NewRule(types.KindLoginRule, RW()),
types.NewRule(types.KindSAMLIdPServiceProvider, RW()),
types.NewRule(types.KindUserGroup, RW()),
types.NewRule(types.KindPlugin, RW()),
types.NewRule(types.KindOktaImportRule, RW()),
types.NewRule(types.KindOktaAssignment, RW()),
types.NewRule(types.KindLock, RW()),
types.NewRule(types.KindIntegration, append(RW(), types.VerbUse)),
types.NewRule(types.KindBilling, RW()),
types.NewRule(types.KindClusterAlert, RW()),
types.NewRule(types.KindAccessList, RW()),
types.NewRule(types.KindNode, RW()),
types.NewRule(types.KindDiscoveryConfig, RW()),
types.NewRule(types.KindSecurityReport, append(RW(), types.VerbUse)),
types.NewRule(types.KindAuditQuery, append(RW(), types.VerbUse)),
types.NewRule(types.KindAccessGraph, RW()),
types.NewRule(types.KindServerInfo, RW()),
types.NewRule(types.KindAccessMonitoringRule, RW()),
types.NewRule(types.KindAppServer, RW()),
types.NewRule(types.KindVnetConfig, RW()),
types.NewRule(types.KindBotInstance, RW()),
types.NewRule(types.KindAccessGraphSettings, RW()),
types.NewRule(types.KindSPIFFEFederation, RW()),
types.NewRule(types.KindNotification, RW()),
types.NewRule(types.KindStaticHostUser, RW()),
types.NewRule(types.KindUserTask, RW()),
types.NewRule(types.KindIdentityCenter, RW()),
types.NewRule(types.KindContact, RW()),
types.NewRule(types.KindWorkloadIdentity, RW()),
types.NewRule(types.KindAutoUpdateVersion, RW()),
types.NewRule(types.KindAutoUpdateConfig, RW()),
types.NewRule(types.KindAutoUpdateAgentRollout, RO()),
types.NewRule(types.KindAutoUpdateAgentReport, RO()),
types.NewRule(types.KindAutoUpdateBotInstanceReport, RO()),
types.NewRule(types.KindGitServer, RW()),
types.NewRule(types.KindWorkloadIdentityX509Revocation, RW()),
types.NewRule(types.KindHealthCheckConfig, RW()),
types.NewRule(types.KindSigstorePolicy, RW()),
types.NewRule(types.KindWorkloadIdentityX509IssuerOverride, RW()),
types.NewRule(types.KindWorkloadIdentityX509IssuerOverrideCSR, RW()),
types.NewRule(types.KindInferenceModel, RW()),
types.NewRule(types.KindInferenceSecret, RW()),
types.NewRule(types.KindInferencePolicy, RW()),
types.NewRule(types.KindClassifier, RW()),
types.NewRule(types.KindRetrievalModel, RW()),
types.NewRule(types.KindClientIPRestriction, RW()),
types.NewRule(access.KindScopedRole, RW()),
types.NewRule(access.KindScopedRoleAssignment, RW()),
types.NewRule(types.KindScopedToken, RW()),
types.NewRule(types.KindWorkloadCluster, RW()),
types.NewRule(types.KindRecordingEncryption, RW()),
types.NewRule(types.KindBeamsConfig, RW()),
},
},
},
}
return role
}
// NewPresetAccessRole creates a role for users who are allowed to initiate
// interactive sessions.
func NewPresetAccessRole() types.Role {
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetAccessRoleName,
Namespace: apidefaults.Namespace,
Description: "Access cluster resources",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
Options: types.RoleOptions{
CertificateFormat: constants.CertificateFormatStandard,
MaxSessionTTL: types.NewDuration(apidefaults.MaxCertDuration),
SSHPortForwarding: &types.SSHPortForwarding{
Remote: &types.SSHRemotePortForwarding{
Enabled: types.NewBoolOption(true),
},
Local: &types.SSHLocalPortForwarding{
Enabled: types.NewBoolOption(true),
},
},
ForwardAgent: types.NewBool(true),
BPF: apidefaults.EnhancedEvents(),
RecordSession: &types.RecordSession{Desktop: types.NewBoolOption(true)},
},
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
Allow: types.RoleConditions{
Namespaces: []string{apidefaults.Namespace},
NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
AppLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
KubernetesLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
WindowsDesktopLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
LinuxDesktopLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseServiceLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseNames: []string{teleport.TraitInternalDBNamesVariable},
DatabaseUsers: []string{teleport.TraitInternalDBUsersVariable},
DatabaseRoles: []string{teleport.TraitInternalDBRolesVariable},
KubernetesResources: []types.KubernetesResource{
{
Kind: types.Wildcard,
Namespace: types.Wildcard,
Name: types.Wildcard,
Verbs: []string{types.Wildcard},
APIGroup: "",
},
},
GitHubPermissions: []types.GitHubPermission{{
Organizations: []string{teleport.TraitInternalGitHubOrgs},
}},
Rules: []types.Rule{
types.NewRule(types.KindEvent, RO()),
{
Resources: []string{types.KindSession},
Verbs: []string{types.VerbRead, types.VerbList},
Where: "contains(session.participants, user.metadata.name)",
},
types.NewRule(types.KindInstance, RO()),
types.NewRule(types.KindClusterMaintenanceConfig, RO()),
},
MCP: &types.MCPPermissions{
Tools: []string{teleport.TraitInternalMCPTools},
},
},
},
}
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
role.SetLogins(types.Allow, []string{teleport.TraitInternalLoginsVariable})
role.SetWindowsLogins(types.Allow, []string{teleport.TraitInternalWindowsLoginsVariable})
role.SetLinuxDesktopLogins(types.Allow, []string{teleport.TraitInternalLinuxDesktopLoginsVariable})
role.SetKubeUsers(types.Allow, []string{teleport.TraitInternalKubeUsersVariable})
role.SetKubeGroups(types.Allow, []string{teleport.TraitInternalKubeGroupsVariable})
role.SetAWSRoleARNs(types.Allow, []string{teleport.TraitInternalAWSRoleARNs})
role.SetAzureIdentities(types.Allow, []string{teleport.TraitInternalAzureIdentities})
role.SetGCPServiceAccounts(types.Allow, []string{teleport.TraitInternalGCPServiceAccounts})
return role
}
// NewPresetAuditorRole returns a new pre-defined role for cluster
// auditor - someone who can review cluster events and replay sessions,
// but can't initiate interactive sessions or modify configuration.
func NewPresetAuditorRole() types.Role {
// IMPORTANT: Before adding new defaults, please make sure that the
// underlying field is supported by the standard role editor UI. This role
// should be editable with a rich UI, without requiring the user to dive into
// YAML.
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetAuditorRoleName,
Namespace: apidefaults.Namespace,
Description: "Review cluster events and replay sessions",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Options: types.RoleOptions{
CertificateFormat: constants.CertificateFormatStandard,
MaxSessionTTL: types.NewDuration(apidefaults.MaxCertDuration),
RecordSession: &types.RecordSession{
Desktop: types.NewBoolOption(false),
},
},
Allow: types.RoleConditions{
Namespaces: []string{apidefaults.Namespace},
Rules: []types.Rule{
types.NewRule(types.KindSession, RO()),
types.NewRule(types.KindEvent, RO()),
types.NewRule(types.KindSessionTracker, RO()),
types.NewRule(types.KindClusterAlert, RO()),
types.NewRule(types.KindInstance, RO()),
types.NewRule(types.KindSecurityReport, append(RO(), types.VerbUse)),
types.NewRule(types.KindAuditQuery, append(RO(), types.VerbUse)),
types.NewRule(types.KindBotInstance, RO()),
types.NewRule(types.KindNotification, RO()),
},
},
},
}
return role
}
// NewPresetReviewerRole returns a new pre-defined role for reviewer. The
// reviewer will be able to review all access requests.
func NewPresetReviewerRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetReviewerRoleName,
Namespace: apidefaults.Namespace,
Description: "Review access requests",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
ReviewRequests: defaultAllowAccessReviewConditions(true)[teleport.PresetReviewerRoleName],
},
},
}
return role
}
// NewPresetRequesterRole returns a new pre-defined role for requester. The
// requester will be able to request all resources.
func NewPresetRequesterRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetRequesterRoleName,
Namespace: apidefaults.Namespace,
Description: "Request all resources",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Request: defaultAllowAccessRequestConditions(true)[teleport.PresetRequesterRoleName],
},
},
}
return role
}
// NewPresetGroupAccessRole returns a new pre-defined role for group access -
// a role used for requesting and reviewing user group access.
func NewPresetGroupAccessRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetGroupAccessRoleName,
Namespace: apidefaults.Namespace,
Description: "Have access to all user groups",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Namespaces: []string{apidefaults.Namespace},
GroupLabels: types.Labels{
types.Wildcard: []string{types.Wildcard},
},
Rules: []types.Rule{
types.NewRule(types.KindUserGroup, RO()),
},
},
},
}
return role
}
// NewPresetDeviceAdminRole returns the preset "device-admin" role, or nil for
// non-Enterprise builds.
// The role is used to administer trusted devices.
func NewPresetDeviceAdminRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
return &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetDeviceAdminRoleName,
Namespace: apidefaults.Namespace,
Description: "Administer trusted devices",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Rules: []types.Rule{
types.NewRule(types.KindDevice, append(RW(), types.VerbCreateEnrollToken, types.VerbEnroll)),
types.NewRule(types.KindMobileDevice, []string{types.VerbCreateEnrollToken}),
},
},
},
}
}
// NewPresetDeviceEnrollRole returns the preset "device-enroll" role, or nil for
// non-Enterprise builds.
// The role is used to grant device enrollment powers to users.
func NewPresetDeviceEnrollRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
return &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetDeviceEnrollRoleName,
Namespace: apidefaults.Namespace,
Description: "Grant permission to enroll trusted devices",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Rules: []types.Rule{
types.NewRule(types.KindDevice, []string{types.VerbEnroll}),
},
},
},
}
}
// NewPresetRequireTrustedDeviceRole returns the preset "require-trusted-device"
// role, or nil for non-Enterprise builds.
// The role is used as a basis for requiring trusted device access to
// resources.
func NewPresetRequireTrustedDeviceRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
return &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetRequireTrustedDeviceRoleName,
Namespace: apidefaults.Namespace,
Description: "Require trusted device to access resources",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Options: types.RoleOptions{
DeviceTrustMode: constants.DeviceTrustModeRequired,
},
Allow: types.RoleConditions{
// All SSH nodes.
Logins: []string{"{{internal.logins}}"},
NodeLabels: types.Labels{
types.Wildcard: []string{types.Wildcard},
},
// All k8s nodes.
KubeGroups: []string{
"{{internal.kubernetes_groups}}",
// Common/example groups.
"system:masters",
"developers",
"viewers",
},
KubernetesLabels: types.Labels{
types.Wildcard: []string{types.Wildcard},
},
// All DB nodes.
DatabaseLabels: types.Labels{
types.Wildcard: []string{types.Wildcard},
},
DatabaseNames: []string{types.Wildcard},
DatabaseUsers: []string{types.Wildcard},
},
},
}
}
// NewPresetWildcardWorkloadIdentityIssuerRole returns a new pre-defined role
// for issuing workload identities.
func NewPresetWildcardWorkloadIdentityIssuerRole() types.Role {
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetWildcardWorkloadIdentityIssuerRoleName,
Namespace: apidefaults.Namespace,
Description: "Issue workload identities",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
WorkloadIdentityLabels: types.Labels{
types.Wildcard: []string{types.Wildcard},
},
Rules: []types.Rule{
types.NewRule(types.KindWorkloadIdentity, RO()),
},
},
},
}
return role
}
// NewPresetAccessPluginRole returns a new pre-defined role for self-hosted
// access request plugins.
func NewPresetAccessPluginRole() types.Role {
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetAccessPluginRoleName,
Namespace: apidefaults.Namespace,
Description: "Default access plugin role",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Rules: []types.Rule{
types.NewRule(types.KindAccessRequest, RO()),
types.NewRule(types.KindAccessPluginData, RW()),
types.NewRule(types.KindAccessMonitoringRule, RO()),
types.NewRule(types.KindAccessList, RO()),
types.NewRule(types.KindRole, RO()),
types.NewRule(types.KindUser, RO()),
types.NewRule(types.KindUserLoginState, RO()),
},
ReviewRequests: &types.AccessReviewConditions{
PreviewAsRoles: []string{
teleport.PresetListAccessRequestResourcesRoleName,
},
},
},
},
}
return role
}
// NewPresetAccessPluginWithReviewRole returns a new pre-defined role for self-hosted
// access request plugins that permits review.
func NewPresetAccessPluginWithReviewRole() types.Role {
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V8,
Metadata: types.Metadata{
Name: teleport.PresetAccessPluginWithReviewRoleName,
Namespace: apidefaults.Namespace,
Description: "Default access plugin with review role",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Rules: []types.Rule{
types.NewRule(types.KindAccessRequest, RO()),
types.NewRule(types.KindAccessPluginData, RW()),
types.NewRule(types.KindAccessMonitoringRule, RO()),
types.NewRule(types.KindAccessList, RO()),
types.NewRule(types.KindRole, RO()),
types.NewRule(types.KindUser, RO()),
types.NewRule(types.KindUserLoginState, RO()),
},
ReviewRequests: &types.AccessReviewConditions{
PreviewAsRoles: []string{
teleport.PresetListAccessRequestResourcesRoleName,
},
SubmitForUsers: []string{"*"},
},
},
},
}
return role
}
// NewPresetListAccessRequestResourcesRole returns a new pre-defined role that
// allows reading access request resources.
func NewPresetListAccessRequestResourcesRole() types.Role {
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetListAccessRequestResourcesRoleName,
Namespace: apidefaults.Namespace,
Description: "Default list access request resources role",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Rules: []types.Rule{
types.NewRule(types.KindNode, RO()),
types.NewRule(types.KindApp, RO()),
types.NewRule(types.KindDatabase, RO()),
types.NewRule(types.KindKubernetesCluster, RO()),
},
// To enable all access plugin features, the role requires read
// access to all of the following resources.
AppLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
GroupLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
KubernetesLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
},
},
}
return role
}
// SystemOktaAccessRoleName is the name of the system role that allows
// access to Okta resources. This will be used by the Okta requester role to
// search for Okta resources.
func NewSystemOktaAccessRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.SystemOktaAccessRoleName,
Namespace: apidefaults.Namespace,
Description: "Request Okta resources",
Labels: map[string]string{
types.TeleportInternalResourceType: types.SystemResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
AppLabels: types.Labels{
types.OriginLabel: []string{types.OriginOkta},
},
GroupLabels: types.Labels{
types.OriginLabel: []string{types.OriginOkta},
},
Rules: []types.Rule{
types.NewRule(types.KindUserGroup, RO()),
},
},
},
}
return role
}
// NewSystemOktaRequesterRole is a system role that allows
// for requesting access to Okta resources. This differs from the requester role
// in that it allows for requesting longer lived access.
func NewSystemOktaRequesterRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.SystemOktaRequesterRoleName,
Namespace: apidefaults.Namespace,
Description: "Request Okta resources",
Labels: map[string]string{
types.TeleportInternalResourceType: types.SystemResource,
types.OriginLabel: types.OriginOkta,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Request: defaultAllowAccessRequestConditions(true)[teleport.SystemOktaRequesterRoleName],
},
},
}
return role
}
// NewSystemIdentityCenterAccessRole creates a role that allows access to AWS
// IdentityCenter resources via Access Requests
func NewSystemIdentityCenterAccessRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
return &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.SystemIdentityCenterAccessRoleName,
Namespace: apidefaults.Namespace,
Description: "Access AWS IAM Identity Center resources",
Labels: map[string]string{
types.TeleportInternalResourceType: types.SystemResource,
// OriginLabel should not be set to AWS Identity center because:
// - identity center is not the one owning this role, this role
// is part of the Teleport system requirements
// - setting the label to a value not support in older agents
// (v16) will cause them to crash.
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
AccountAssignments: defaultAllowAccountAssignments(true)[teleport.SystemIdentityCenterAccessRoleName],
},
},
}
}
// NewPresetTerraformProviderRole returns a new pre-defined role for the Teleport Terraform provider.
// This role can edit any Terraform-supported resource.
func NewPresetTerraformProviderRole() types.Role {
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: teleport.PresetTerraformProviderRoleName,
Namespace: apidefaults.Namespace,
Description: "Default Terraform provider role",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
// In Teleport, you can only see what you have access to. To be able to reconcile
// Apps, Databases, Dynamic Windows Desktops, and Nodes, Terraform must be able to
// access them all.
// For Databases and Nodes, Terraform cannot actually access them because it has no
// Login/user set.
AppLabels: map[string]apiutils.Strings{types.Wildcard: []string{types.Wildcard}},
DatabaseLabels: map[string]apiutils.Strings{types.Wildcard: []string{types.Wildcard}},
KubernetesLabels: map[string]apiutils.Strings{types.Wildcard: []string{types.Wildcard}},
NodeLabels: map[string]apiutils.Strings{types.Wildcard: []string{types.Wildcard}},
WindowsDesktopLabels: map[string]apiutils.Strings{types.Wildcard: []string{types.Wildcard}},
// Every resource currently supported by the Terraform provider.
Rules: []types.Rule{
// You must add new resources as separate rules for the
// default rule addition logic to work properly.
types.NewRule(types.KindAccessList, RW()),
types.NewRule(types.KindApp, RW()),
types.NewRule(types.KindClusterAuthPreference, RW()),
types.NewRule(types.KindClusterMaintenanceConfig, RW()),
types.NewRule(types.KindClusterNetworkingConfig, RW()),
types.NewRule(types.KindDatabase, RW()),
types.NewRule(types.KindDevice, RW()),
types.NewRule(types.KindDiscoveryConfig, RW()),
types.NewRule(types.KindGithub, RW()),
types.NewRule(types.KindKubernetesCluster, RW()),
types.NewRule(types.KindLock, RW()),
types.NewRule(types.KindLoginRule, RW()),
types.NewRule(types.KindNode, RW()),
types.NewRule(types.KindOIDC, RW()),
types.NewRule(types.KindOktaImportRule, RW()),
types.NewRule(types.KindRole, RW()),
types.NewRule(types.KindSAML, RW()),
types.NewRule(types.KindSessionRecordingConfig, RW()),
types.NewRule(types.KindToken, RW()),
types.NewRule(types.KindTrustedCluster, RW()),
types.NewRule(types.KindUIConfig, RW()),
types.NewRule(types.KindUser, RW()),
types.NewRule(types.KindBot, RW()),
types.NewRule(types.KindInstaller, RW()),
types.NewRule(types.KindAccessMonitoringRule, RW()),
types.NewRule(types.KindDynamicWindowsDesktop, RW()),
types.NewRule(types.KindStaticHostUser, RW()),
types.NewRule(types.KindWorkloadIdentity, RW()),
types.NewRule(types.KindGitServer, RW()),
types.NewRule(types.KindAutoUpdateConfig, RW()),
types.NewRule(types.KindAutoUpdateVersion, RW()),
types.NewRule(types.KindHealthCheckConfig, RW()),
types.NewRule(types.KindVnetConfig, RW()),
types.NewRule(types.KindIntegration, RW()),
types.NewRule(types.KindInferenceModel, RW()),
types.NewRule(types.KindInferenceSecret, RW()),
types.NewRule(types.KindInferencePolicy, RW()),
types.NewRule(types.KindClassifier, RW()),
types.NewRule(types.KindRetrievalModel, RW()),
types.NewRule(types.KindSAMLIdPServiceProvider, RW()),
types.NewRule(types.KindScopedToken, RW()),
types.NewRule(access.KindScopedRole, RW()),
types.NewRule(access.KindScopedRoleAssignment, RW()),
types.NewRule(types.KindDatabaseObjectImportRule, RW()),
types.NewRule(types.KindBeamsConfig, RW()),
},
},
},
}
return role
}
// NewPresetMCPUserRole returns a new pre-defined role for accessing MCP
// servers.
func NewPresetMCPUserRole() types.Role {
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V8,
Metadata: types.Metadata{
Name: teleport.PresetMCPUserRoleName,
Namespace: apidefaults.Namespace,
Description: "Access to MCP servers",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
AppLabels: map[string]apiutils.Strings{
types.AppSubKindLabel: []string{types.SubKindMCP},
},
MCP: &types.MCPPermissions{
Tools: []string{types.Wildcard},
},
},
},
}
return role
}
// NewPresetBeamUserRole returns a new pre-defined role for accessing your own
// beam resources.
func NewPresetBeamUserRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
allowLLMApps := fmt.Sprintf(`labels[%q] == %q`, types.BeamAppTypeLabel, types.SubKindLLM)
allowBeamApps := fmt.Sprintf(`labels[%q] == user.metadata.name`, types.BeamOwnerLabel)
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V8,
Metadata: types.Metadata{
Name: teleport.PresetBeamUserRoleName,
Namespace: apidefaults.Namespace,
Description: "Use the Beams feature",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Logins: []string{types.BeamsLogin},
AppLabelsExpression: strings.Join([]string{allowLLMApps, allowBeamApps}, " || "),
NodeLabels: types.Labels{
types.BeamOwnerLabel: {"{{user.metadata.name}}"},
},
BeamLabels: types.Labels{
types.BeamOwnerLabel: {"{{user.metadata.name}}"},
},
Rules: []types.Rule{
{
Resources: []string{types.KindBeam},
Verbs: []string{types.Wildcard},
},
types.NewRule(types.KindBeamsConfig, RO()),
},
},
},
}
return role
}
// NewPresetBeamAdminRole returns a new pre-defined role for administering beams
// belonging to other users.
func NewPresetBeamAdminRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V8,
Metadata: types.Metadata{
Name: teleport.PresetBeamAdminRoleName,
Namespace: apidefaults.Namespace,
Description: "Administer beams belonging to other users",
Labels: map[string]string{
types.TeleportInternalResourceType: types.PresetResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
BeamLabels: types.Labels{
types.BeamOwnerLabel: {types.Wildcard},
},
Rules: []types.Rule{
{
Resources: []string{types.KindBeam},
Verbs: []string{types.Wildcard},
},
types.NewRule(types.KindBeamsConfig, RW()),
},
},
},
}
return role
}
// NewSystemBeamRole returns a new pre-defined role for the beam to issue itself
// credentials.
func NewSystemBeamRole(buildType string) types.Role {
if buildType != modules.BuildEnterprise {
return nil
}
// Only allow the bot to generate host certificates for the Beam's OpenSSH
// server. tbot explicitly sends a blank HostID, HostName and the Node role.
hostCertConstraints := strings.Join([]string{
fmt.Sprintf(`contains_all(user.spec.traits[%q], host_cert.principals)`, types.BeamIDLabel),
`host_cert.host_id == ""`,
`host_cert.node_name == ""`,
}, " && ")
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V8,
Metadata: types.Metadata{
Name: teleport.SystemBeamRoleName,
Namespace: apidefaults.Namespace,
Description: "Used by a beam to issue itself credentials",
Labels: map[string]string{
types.TeleportInternalResourceType: types.SystemResource,
},
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{
Rules: []types.Rule{
{
Resources: []string{types.KindHostCert},
Verbs: []string{types.VerbCreate},
Where: hostCertConstraints,
},
{
Resources: []string{types.KindWorkloadIdentity},
Verbs: []string{
types.VerbList,
types.VerbRead,
},
},
},
WorkloadIdentityLabels: types.Labels{
types.BeamIDLabel: []string{fmt.Sprintf(`{{external[%q]}}`, types.BeamIDLabel)},
},
},
},
}
return role
}
// VirtualDefaultHealthCheckConfigDB returns a health_check_config enabling
// health checks for all databases resources, and is intended to be used as a
// virtual default resource. Its name is "default" for historical reasons.
func VirtualDefaultHealthCheckConfigDB() *healthcheckconfigv1.HealthCheckConfig {
return healthcheckconfigv1.HealthCheckConfig_builder{
Kind: types.KindHealthCheckConfig,
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: teleport.VirtualDefaultHealthCheckConfigDBName,
Description: "Enables health checks for all databases by default",
// this revision MUST be changed every time we change the contents
// of the preset so that conditional updates can check against it
Revision: "af391615-1e42-4237-aa2b-155e6abbd41a",
}.Build(),
Spec: healthcheckconfigv1.HealthCheckConfigSpec_builder{
Match: healthcheckconfigv1.Matcher_builder{
// match all databases
DbLabels: []*labelv1.Label{labelv1.Label_builder{
Name: types.Wildcard,
Values: []string{types.Wildcard},
}.Build()},
}.Build(),
}.Build(),
}.Build()
}
// VirtualDefaultHealthCheckConfigKube returns a health_check_config enabling
// health checks for all Kubernetes resources. It's intended to be used as a
// virtual default resource.
func VirtualDefaultHealthCheckConfigKube() *healthcheckconfigv1.HealthCheckConfig {
return healthcheckconfigv1.HealthCheckConfig_builder{
Kind: types.KindHealthCheckConfig,
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: teleport.VirtualDefaultHealthCheckConfigKubeName,
Description: "Enables health checks for all Kubernetes clusters by default.",
// this revision MUST be changed every time we change the contents
// of the preset so that conditional updates can check against it
Revision: "d796f007-e60c-4747-8dde-f479aff6b743",
}.Build(),
Spec: healthcheckconfigv1.HealthCheckConfigSpec_builder{
Match: healthcheckconfigv1.Matcher_builder{
// match all kubernetes clusters
KubernetesLabels: []*labelv1.Label{labelv1.Label_builder{
Name: types.Wildcard,
Values: []string{types.Wildcard},
}.Build()},
}.Build(),
}.Build(),
}.Build()
}
// bootstrapRoleMetadataLabels are metadata labels that will be applied to each role.
// These are intended to add labels for older roles that didn't previously have them.
func bootstrapRoleMetadataLabels() map[string]map[string]string {
return map[string]map[string]string{
teleport.PresetAccessRoleName: {
types.TeleportInternalResourceType: types.PresetResource,
},
teleport.PresetEditorRoleName: {
types.TeleportInternalResourceType: types.PresetResource,
},
teleport.PresetAuditorRoleName: {
types.TeleportInternalResourceType: types.PresetResource,
},
teleport.SystemOktaRequesterRoleName: {
types.TeleportInternalResourceType: types.SystemResource,
types.OriginLabel: types.OriginOkta,
},
// We unset the OriginLabel on the system AWS IC role because this value
// was not supported on v16 agents and this crashes them.
teleport.SystemIdentityCenterAccessRoleName: {
types.TeleportInternalResourceType: types.SystemResource,
},
// These roles are intentionally not added here as there may be existing
// customer defined roles that have these labels:
// group-access, reviewer, requester, mcp-user
}
}
// defaultAllowRules has the Allow rules that should be set as default when
// they were not explicitly defined. This is used to update the current cluster
// roles when deploying a new resource. It will also update all existing roles
// on auth server restart. Rules defined in preset template should be
// exactly the same rule when added here.
func defaultAllowRules(buildType string) map[string][]types.Rule {
roles := []types.Role{
NewPresetAuditorRole(),
NewPresetEditorRole(),
NewPresetAccessRole(),
NewPresetTerraformProviderRole(),
NewPresetAccessPluginRole(),
NewPresetAccessPluginWithReviewRole(),
NewPresetListAccessRequestResourcesRole(),
NewPresetDeviceAdminRole(buildType),
NewPresetBeamUserRole(buildType),
NewPresetBeamAdminRole(buildType),
}
allowRules := make(map[string][]types.Rule, len(roles))
for _, role := range roles {
if role == nil {
continue
}
allowRules[role.GetName()] = role.GetRules(types.Allow)
}
return allowRules
}
// defaultAllowLabels has the Allow labels that should be set as default when they were not explicitly defined.
// This is used to update existing builtin preset roles with new permissions during cluster upgrades.
// The following Labels are supported:
// - AppLabels
// - DatabaseServiceLabels (db_service_labels)
// - GroupLabels
func defaultAllowLabels(enterprise bool) map[string]types.RoleConditions {
wildcardLabels := types.Labels{types.Wildcard: []string{types.Wildcard}}
conditions := map[string]types.RoleConditions{
teleport.PresetAccessRoleName: {
DatabaseServiceLabels: wildcardLabels,
DatabaseRoles: []string{teleport.TraitInternalDBRolesVariable},
},
teleport.PresetTerraformProviderRoleName: {
AppLabels: wildcardLabels,
DatabaseLabels: wildcardLabels,
NodeLabels: wildcardLabels,
KubernetesLabels: wildcardLabels,
WindowsDesktopLabels: wildcardLabels,
},
teleport.PresetListAccessRequestResourcesRoleName: {
AppLabels: wildcardLabels,
DatabaseLabels: wildcardLabels,
GroupLabels: wildcardLabels,
KubernetesLabels: wildcardLabels,
NodeLabels: wildcardLabels,
},
}
if enterprise {
conditions[teleport.SystemOktaAccessRoleName] = types.RoleConditions{
AppLabels: types.Labels{types.OriginLabel: []string{types.OriginOkta}},
GroupLabels: types.Labels{types.OriginLabel: []string{types.OriginOkta}},
}
}
return conditions
}
// defaultAllowAccessRequestConditions has the access request conditions that should be set as default when they were
// not explicitly defined.
func defaultAllowAccessRequestConditions(enterprise bool) map[string]*types.AccessRequestConditions {
if enterprise {
return map[string]*types.AccessRequestConditions{
teleport.PresetRequesterRoleName: {
SearchAsRoles: []string{
teleport.PresetAccessRoleName,
teleport.PresetGroupAccessRoleName,
teleport.SystemIdentityCenterAccessRoleName,
},
},
teleport.SystemOktaRequesterRoleName: {
SearchAsRoles: []string{
teleport.SystemOktaAccessRoleName,
},
MaxDuration: types.NewDuration(MaxAccessDuration),
},
}
}
return map[string]*types.AccessRequestConditions{}
}
// defaultAllowAccessReviewConditions has the access review conditions that should be set as default when they were
// not explicitly defined.
func defaultAllowAccessReviewConditions(enterprise bool) map[string]*types.AccessReviewConditions {
if enterprise {
return map[string]*types.AccessReviewConditions{
teleport.PresetReviewerRoleName: {
PreviewAsRoles: []string{
teleport.PresetAccessRoleName,
teleport.PresetGroupAccessRoleName,
teleport.SystemIdentityCenterAccessRoleName,
},
Roles: []string{
teleport.PresetAccessRoleName,
teleport.PresetGroupAccessRoleName,
teleport.SystemIdentityCenterAccessRoleName,
},
},
}
}
return map[string]*types.AccessReviewConditions{}
}
func defaultAllowAccountAssignments(enterprise bool) map[string][]types.IdentityCenterAccountAssignment {
if enterprise {
return map[string][]types.IdentityCenterAccountAssignment{
teleport.SystemIdentityCenterAccessRoleName: {
{
Account: types.Wildcard,
PermissionSet: types.Wildcard,
},
},
}
}
return map[string][]types.IdentityCenterAccountAssignment{}
}
// AddRoleDefaults adds default role attributes to a preset role.
// Only attributes whose resources are not already defined (either allowing or denying) are added.
func AddRoleDefaults(ctx context.Context, buildType string, role types.Role) (types.Role, error) {
changed := false
oldLabels := role.GetAllLabels()
// Role labels
defaultRoleLabels, ok := bootstrapRoleMetadataLabels()[role.GetName()]
if ok {
metadata := role.GetMetadata()
if metadata.Labels == nil {
metadata.Labels = make(map[string]string, len(defaultRoleLabels))
}
for label, value := range defaultRoleLabels {
if _, ok := metadata.Labels[label]; !ok {
metadata.Labels[label] = value
changed = true
}
}
if changed {
role.SetMetadata(metadata)
}
}
labels := role.GetMetadata().Labels
// We're specifically checking the old labels version of the Okta requester role here
// because we're bootstrapping new labels onto the role above. By checking the old labels,
// we can be assured that we're looking at the role as it existed before bootstrapping. If
// the role was user-created, then this won't have the internal-resource type attached,
// and we'll skip the rest of adding in default values.
if role.GetName() == teleport.SystemOktaRequesterRoleName {
labels = oldLabels
}
// Check if the role has a TeleportInternalResourceType attached. We do this after setting the role metadata
// labels because we set the role metadata labels for roles that have been well established (access,
// editor, auditor) that may not already have this label set, but we don't set it for newer roles
// (group-access, reviewer, requester, mcp-user) that may have customer definitions.
resourceType := labels[types.TeleportInternalResourceType]
if resourceType != types.PresetResource && resourceType != types.SystemResource {
return nil, trace.AlreadyExists("not modifying user created role")
}
// Resource Rules
defaultRules, ok := defaultAllowRules(buildType)[role.GetName()]
if ok {
existingRules := append(role.GetRules(types.Allow), role.GetRules(types.Deny)...)
for _, defaultRule := range defaultRules {
if resourceBelongsToRules(existingRules, defaultRule.Resources) {
continue
}
slog.DebugContext(ctx, "Adding default allow rule to role",
"rule", defaultRule,
"role", role.GetName(),
)
rules := role.GetRules(types.Allow)
rules = append(rules, defaultRule)
role.SetRules(types.Allow, rules)
changed = true
}
}
enterprise := buildType == modules.BuildEnterprise
// Labels
defaultLabels, ok := defaultAllowLabels(enterprise)[role.GetName()]
if ok {
for _, kind := range []string{
types.KindApp,
types.KindDatabase,
types.KindDatabaseService,
types.KindNode,
types.KindUserGroup,
types.KindWindowsDesktop,
types.KindKubernetesCluster,
} {
var labels types.Labels
switch kind {
case types.KindApp:
labels = defaultLabels.AppLabels
case types.KindDatabase:
labels = defaultLabels.DatabaseLabels
case types.KindDatabaseService:
labels = defaultLabels.DatabaseServiceLabels
case types.KindNode:
labels = defaultLabels.NodeLabels
case types.KindUserGroup:
labels = defaultLabels.GroupLabels
case types.KindWindowsDesktop:
labels = defaultLabels.WindowsDesktopLabels
case types.KindKubernetesCluster:
labels = defaultLabels.KubernetesLabels
}
labelsUpdated, err := updateAllowLabels(role, kind, labels)
if err != nil {
return nil, trace.Wrap(err)
}
changed = changed || labelsUpdated
}
if len(defaultLabels.DatabaseRoles) > 0 && len(role.GetDatabaseRoles(types.Allow)) == 0 {
role.SetDatabaseRoles(types.Allow, defaultLabels.DatabaseRoles)
changed = true
}
}
if roleUpdated := applyAccessRequestConditionDefaults(role, enterprise); roleUpdated {
changed = true
}
if roleUpdated := applyAccessReviewConditionDefaults(role, enterprise); roleUpdated {
changed = true
}
if len(role.GetIdentityCenterAccountAssignments(types.Allow)) == 0 {
assignments := defaultAllowAccountAssignments(enterprise)[role.GetName()]
if assignments != nil {
role.SetIdentityCenterAccountAssignments(types.Allow, assignments)
changed = true
}
}
// GitHub permissions.
if len(role.GetGitHubPermissions(types.Allow)) == 0 {
if githubOrgs := defaultGitHubOrgs()[role.GetName()]; len(githubOrgs) > 0 {
role.SetGitHubPermissions(types.Allow, []types.GitHubPermission{{
Organizations: githubOrgs,
}})
changed = true
}
}
if role.GetMCPPermissions(types.Allow) == nil {
if mcpTools := defaultMCPTools()[role.GetName()]; len(mcpTools) > 0 {
role.SetMCPPermissions(types.Allow, &types.MCPPermissions{
Tools: mcpTools,
})
changed = true
}
}
if !changed {
return nil, trace.AlreadyExists("no change")
}
return role, nil
}
func mergeStrings(dst, src []string) (merged []string, changed bool) {
items := set.New[string](dst...)
items.Add(src...)
if len(items) == len(dst) {
return dst, false
}
dst = items.Elements()
slices.Sort(dst)
return dst, true
}
func applyAccessRequestConditionDefaults(role types.Role, enterprise bool) bool {
defaults := defaultAllowAccessRequestConditions(enterprise)[role.GetName()]
if defaults == nil {
return false
}
target := role.GetAccessRequestConditions(types.Allow)
changed := false
if target.IsEmpty() {
target = *defaults
changed = true
} else {
var rolesUpdated bool
target.Roles, rolesUpdated = mergeStrings(target.Roles, defaults.Roles)
changed = changed || rolesUpdated
target.SearchAsRoles, rolesUpdated = mergeStrings(target.SearchAsRoles, defaults.SearchAsRoles)
changed = changed || rolesUpdated
}
if changed {
role.SetAccessRequestConditions(types.Allow, target)
}
return changed
}
func applyAccessReviewConditionDefaults(role types.Role, enterprise bool) bool {
defaults := defaultAllowAccessReviewConditions(enterprise)[role.GetName()]
if defaults == nil {
return false
}
target := role.GetAccessReviewConditions(types.Allow)
changed := false
if target.IsEmpty() {
target = *defaults
changed = true
} else {
var rolesUpdated bool
target.Roles, rolesUpdated = mergeStrings(target.Roles, defaults.Roles)
changed = changed || rolesUpdated
target.PreviewAsRoles, rolesUpdated = mergeStrings(target.PreviewAsRoles, defaults.PreviewAsRoles)
changed = changed || rolesUpdated
}
if changed {
role.SetAccessReviewConditions(types.Allow, target)
}
return changed
}
func labelMatchersUnset(role types.Role, kind string) (bool, error) {
for _, cond := range []types.RoleConditionType{types.Allow, types.Deny} {
labelMatchers, err := role.GetLabelMatchers(cond, kind)
if err != nil {
return false, trace.Wrap(err)
}
if !labelMatchers.Empty() {
return false, nil
}
}
return true, nil
}
func resourceBelongsToRules(rules []types.Rule, resources []string) bool {
for _, rule := range rules {
for _, ruleResource := range rule.Resources {
if slices.Contains(resources, ruleResource) {
return true
}
}
}
return false
}
func updateAllowLabels(role types.Role, kind string, defaultLabels types.Labels) (bool, error) {
var changed bool
if unset, err := labelMatchersUnset(role, kind); err != nil {
return false, trace.Wrap(err)
} else if unset && len(defaultLabels) > 0 {
role.SetLabelMatchers(types.Allow, kind, types.LabelMatchers{
Labels: defaultLabels,
})
changed = true
}
return changed, nil
}
func defaultGitHubOrgs() map[string][]string {
return map[string][]string{
teleport.PresetAccessRoleName: {teleport.TraitInternalGitHubOrgs},
}
}
func defaultMCPTools() map[string][]string {
return map[string][]string{
teleport.PresetAccessRoleName: {teleport.TraitInternalMCPTools},
}
}
/*
* 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"
// ProcessForkedContext adds a flag to the context to indicate the Teleport
// process has running forked child(ren).
func ProcessForkedContext(parent context.Context) context.Context {
return addFlagToContext[processForkedFlag](parent)
}
// HasProcessForked returns true if the Teleport process has running forked
// child(ren).
func HasProcessForked(ctx context.Context) bool {
return getFlagFromContext[processForkedFlag](ctx)
}
// ShouldDeleteServerHeartbeatsOnShutdown checks whether server heartbeats
// should be deleted based on the process shutdown context.
func ShouldDeleteServerHeartbeatsOnShutdown(ctx context.Context) bool {
switch {
// A child process can be forked to upgrade the Teleport binary. The child
// will take over the heartbeats so do NOT delete them in that case. In
// worst case scenarios if the child fails to register new heartbeats, the
// old ones will get deleted automatically upon expiry.
case HasProcessForked(ctx):
return false
default:
return true
}
}
func addFlagToContext[FlagType any](parent context.Context) context.Context {
return context.WithValue(parent, (*FlagType)(nil), (*FlagType)(nil))
}
func getFlagFromContext[FlagType any](ctx context.Context) bool {
_, ok := ctx.Value((*FlagType)(nil)).(*FlagType)
return ok
}
type processForkedFlag struct{}
/*
* 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"
"strings"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/utils"
)
// Provisioner governs adding new nodes to the cluster
type Provisioner interface {
// UpsertToken adds provisioning tokens for the auth server
UpsertToken(ctx context.Context, token types.ProvisionToken) error
// CreateToken adds provisioning tokens for the auth server
CreateToken(ctx context.Context, token types.ProvisionToken) error
// GetToken finds and returns token by id
GetToken(ctx context.Context, token string) (types.ProvisionToken, error)
// DeleteToken deletes provisioning token
// Imlementations must guarantee that this returns trace.NotFound error if the token doesn't exist
DeleteToken(ctx context.Context, token string) error
// PatchToken performs a conditional update on the named token using
// `updateFn`, retrying internally if a comparison failure occurs.
PatchToken(
ctx context.Context,
token string,
updateFn func(types.ProvisionToken) (types.ProvisionToken, error),
) (types.ProvisionToken, error)
// ListProvisionTokens retrieves a paginated list of provision tokens.
ListProvisionTokens(ctx context.Context, pageSize int, pageToken string, anyRoles types.SystemRoles, botName string) ([]types.ProvisionToken, string, error)
}
// ProvisionerInternal extends the Provisioner interface with auth-specific internal methods.
type ProvisionerInternal interface {
Provisioner
// AppendPutProvisionTokenActions adds conditional actions to an atomic write
// to create or update a provision token.
AppendPutProvisionTokenActions(
actions []backend.ConditionalAction,
token types.ProvisionToken,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteProvisionTokenActions adds conditional actions to an atomic
// write to delete a provision token.
AppendDeleteProvisionTokenActions(
actions []backend.ConditionalAction,
token string,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
// MustCreateProvisionToken returns a new valid provision token
// or panics, used in tests
func MustCreateProvisionToken(token string, roles types.SystemRoles, expires time.Time) types.ProvisionToken {
t, err := types.NewProvisionToken(token, roles, expires)
if err != nil {
panic(err)
}
return t
}
// UnmarshalProvisionToken unmarshals the ProvisionToken resource from JSON.
func UnmarshalProvisionToken(data []byte, opts ...MarshalOption) (types.ProvisionToken, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing provision token data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
err = utils.FastUnmarshal(data, &h)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case "":
var p types.ProvisionTokenV1
err := utils.FastUnmarshal(data, &p)
if err != nil {
return nil, trace.Wrap(err)
}
v2 := p.V2()
if cfg.Revision != "" {
v2.SetRevision(cfg.Revision)
}
return v2, nil
case types.V2:
var p types.ProvisionTokenV2
if err := utils.FastUnmarshal(data, &p); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := p.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
p.SetRevision(cfg.Revision)
}
return &p, nil
}
return nil, trace.BadParameter("server resource version %v is not supported", h.Version)
}
// strongValidateProvisionTokenWithDefaults checks if the provision token is valid and sets defaults if necessary..
func strongValidateProvisionTokenWithDefaults(token *types.ProvisionTokenV2) error {
if err := token.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
// for now there are no additional, on-write validations for token types other than kubernetes
if token.GetJoinMethod() != types.JoinMethodKubernetes {
return nil
}
kube := token.GetKubernetes()
if kube == nil {
// technically should never happen since CheckAndSetDefaults() performs a similar check,
// but we'll be defensive just in case
return trace.BadParameter("allow: at least one rule must be set")
}
for i, rule := range kube.Allow {
serviceAccountSet := rule.ServiceAccount != ""
// validation for empty namespace and account was added much later than the rest of the validations
// in CheckAndSetDefaults(), so we only enforce them when marshaling a token rather than when unmarshaling
if serviceAccountSet {
namespace, account, _ := strings.Cut(rule.ServiceAccount, ":")
if namespace == "" || account == "" {
return trace.BadParameter(
`allow[%d].service_account: name of service account should be in format "namespace:service_account", got %q instead`,
i,
rule.ServiceAccount,
)
}
}
}
return nil
}
// MarshalProvisionToken marshals the ProvisionToken resource to JSON.
func MarshalProvisionToken(provisionToken types.ProvisionToken, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch provisionToken := provisionToken.(type) {
case *types.ProvisionTokenV2:
if err := strongValidateProvisionTokenWithDefaults(provisionToken); err != nil {
return nil, trace.Wrap(err)
}
provisionToken = maybeResetProtoRevision(cfg.PreserveRevision, provisionToken)
if cfg.GetVersion() == types.V1 {
return utils.FastMarshal(provisionToken.V1())
}
return utils.FastMarshal(provisionToken)
default:
return nil, trace.BadParameter("unrecognized provision token version %T", provisionToken)
}
}
// CloneProvisionToken returns a deep copy of the given provision token, per
// `apiutils.CloneProtoMsg()`. Fields in the clone may be modified without
// affecting the original. Only V2 is supported.
func CloneProvisionToken(provisionToken types.ProvisionToken) (types.ProvisionToken, error) {
switch provisionToken := provisionToken.(type) {
case *types.ProvisionTokenV2:
clone := apiutils.CloneProtoMsg(provisionToken)
return clone, nil
default:
return nil, trace.BadParameter("cannot clone unsupported provision token version %T", provisionToken)
}
}
/*
* 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"
"log/slog"
"sync"
"sync/atomic"
"time"
"github.com/gravitational/trace"
"github.com/prometheus/client_golang/prometheus"
"golang.org/x/sync/errgroup"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/observability/metrics"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// Matcher is used by reconciler to match resources.
type Matcher[T any] func(T) bool
// GenericReconcilerConfig is the resource reconciler configuration that allows
// any type implementing comparable to be a key.
type GenericReconcilerConfig[K comparable, T any] struct {
// Matcher is used to match resources.
Matcher Matcher[T]
// GetCurrentResources returns currently registered resources. Note that the
// map keys must be consistent across the current and new resources.
GetCurrentResources func() map[K]T
// GetNewResources returns resources to compare current resources against.
// Note that the map keys must be consistent across the current and new
// resources.
GetNewResources func() map[K]T
// Compare allows custom comparators without having to implement IsEqual.
// Defaults to `CompareResources[T]` if not specified.
CompareResources func(T, T) int
// OnCreate is called when a new resource is detected.
OnCreate func(context.Context, T) error
// OnUpdate is called when an existing resource is updated.
OnUpdate func(ctx context.Context, new, old T) error
// OnDelete is called when an existing resource is deleted.
OnDelete func(context.Context, T) error
// Logger emits log messages.
Logger *slog.Logger
// Metrics is an optional ReconcilerMetrics created by the caller.
// The caller is responsible for registering the metrics.
// Metrics can be nil, in this case the generic reconciler will generate its
// own metrics, which won't be registered.
// Passing a metrics struct might look like a cumbersome API but we have 2 challenges:
// - some parts of Teleport are using one-shot reconcilers. Registering
// metrics on every run would fail and we would lose the past reconciliation
// data.
// - we have many reconcilers in Teleport and making the caller create the
// metrics beforehand allows them to specify the metric subsystem.
Metrics *ReconcilerMetrics
// AllowOriginChanges is a flag that allows the reconciler to change the
// origin value of a reconciled resource. By default, origin changes are
// disallowed to enforce segregation between of resources from different
// sources.
AllowOriginChanges bool
// Concurrency sets the number of goroutines used to process resources
// during reconciliation. When set to 0 or 1, resources are processed
// sequentially. When set to a value greater than 1, resources are
// processed concurrently using up to that many goroutines.
// The OnCreate, OnUpdate, OnDelete, Matcher, and CompareResources
// callbacks must be safe for concurrent use when Concurrency > 1.
Concurrency int
}
// CheckAndSetDefaults validates the reconciler configuration and sets defaults.
func (c *GenericReconcilerConfig[K, T]) CheckAndSetDefaults() error {
if c.Matcher == nil {
return trace.BadParameter("missing reconciler Matcher")
}
if c.GetCurrentResources == nil {
return trace.BadParameter("missing reconciler GetCurrentResources")
}
if c.GetNewResources == nil {
return trace.BadParameter("missing reconciler GetNewResources")
}
if c.OnCreate == nil {
return trace.BadParameter("missing reconciler OnCreate")
}
if c.OnUpdate == nil {
return trace.BadParameter("missing reconciler OnUpdate")
}
if c.OnDelete == nil {
return trace.BadParameter("missing reconciler OnDelete")
}
if c.CompareResources == nil {
return trace.BadParameter("missing reconciler CompareResources")
}
if c.Logger == nil {
c.Logger = slog.With(teleport.ComponentKey, "reconciler")
}
if c.Concurrency < 1 {
c.Concurrency = 1
}
if c.Metrics == nil {
var err error
// If we are not given metrics, we create our own so we don't
// panic when trying to increment/observe.
c.Metrics, err = NewReconcilerMetrics(metrics.NoopRegistry().Wrap("unknown"))
if err != nil {
return trace.Wrap(err)
}
}
return nil
}
// ReconcilerMetrics is a set of metrics that the reconciler will update during
// its reconciliation cycle.
type ReconcilerMetrics struct {
reconciliationTotal *prometheus.CounterVec
reconciliationDuration *prometheus.HistogramVec
}
const (
metricLabelResult = "result"
metricLabelResultSuccess = "success"
metricLabelResultError = "error"
metricLabelResultNoop = "noop"
metricLabelOperation = "operation"
metricLabelOperationCreate = "create"
metricLabelOperationUpdate = "update"
metricLabelOperationDelete = "delete"
metricLabelKind = "kind"
)
// NewReconcilerMetrics creates subsystem-scoped metrics for the reconciler.
// The caller is responsible for registering them into an appropriate registry.
// The same ReconcilerMetrics can be used across different reconcilers.
// The metrics subsystem cannot be empty.
func NewReconcilerMetrics(reg *metrics.Registry) (*ReconcilerMetrics, error) {
if reg == nil {
return nil, trace.BadParameter("missing metrics registry (this is a bug)")
}
if reg.Subsystem() == "" {
return nil, trace.BadParameter("missing metrics subsystem (this is a bug)")
}
return &ReconcilerMetrics{
reconciliationTotal: prometheus.NewCounterVec(prometheus.CounterOpts{
Namespace: reg.Namespace(),
Subsystem: reg.Subsystem(),
Name: "reconciliation_total",
Help: "Total number of individual resource reconciliations.",
}, []string{metricLabelKind, metricLabelOperation, metricLabelResult}),
reconciliationDuration: prometheus.NewHistogramVec(prometheus.HistogramOpts{
Namespace: reg.Namespace(),
Subsystem: reg.Subsystem(),
Name: "reconciliation_duration_seconds",
Help: "The duration of individual resource reconciliation in seconds.",
}, []string{metricLabelKind, metricLabelOperation}),
}, nil
}
// Register metrics in the specified [prometheus.Registerer], returns an error
// if any metric fails, but still tries to register every metric before returning.
func (m *ReconcilerMetrics) Register(r prometheus.Registerer) error {
return trace.NewAggregate(
r.Register(m.reconciliationTotal),
r.Register(m.reconciliationDuration),
)
}
// NewGenericReconciler creates a new GenericReconciler with provided configuration.
func NewGenericReconciler[K comparable, T any](cfg GenericReconcilerConfig[K, T]) (*GenericReconciler[K, T], error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &GenericReconciler[K, T]{
cfg: cfg,
logger: cfg.Logger,
metrics: cfg.Metrics,
stats: &reconcileStats{},
}, nil
}
// GenericReconciler reconciles currently registered resources with new
// resources and creates/updates/deletes them appropriately.
//
// It's used in combination with watchers by agents (app, database, desktop)
// to enable dynamically registered resources.
type GenericReconciler[K comparable, T any] struct {
cfg GenericReconcilerConfig[K, T]
logger *slog.Logger
metrics *ReconcilerMetrics
stats *reconcileStats
}
// reconcileStats tracks the number of resources created, updated, and deleted
// during a reconciliation cycle.
type reconcileStats struct {
created atomic.Int64
updated atomic.Int64
deleted atomic.Int64
}
func (s *reconcileStats) reset() {
s.created.Store(0)
s.updated.Store(0)
s.deleted.Store(0)
}
func (s *reconcileStats) hasChanges() bool {
return s.created.Load() > 0 || s.updated.Load() > 0 || s.deleted.Load() > 0
}
// LogValue implements [slog.LogValuer].
func (s *reconcileStats) LogValue() slog.Value {
return slog.GroupValue(
slog.Int64("created", s.created.Load()),
slog.Int64("updated", s.updated.Load()),
slog.Int64("deleted", s.deleted.Load()),
)
}
// onCreate wraps the OnCreate callback with metrics and stats observation.
func (r *GenericReconciler[K, T]) onCreate(ctx context.Context, kind string, newT T) error {
start := time.Now()
err := r.cfg.OnCreate(ctx, newT)
if err == nil {
r.stats.created.Add(1)
}
r.observeMetrics(kind, metricLabelOperationCreate, start, err)
return trace.Wrap(err)
}
// onUpdate wraps the OnUpdate callback with metrics and stats observation.
func (r *GenericReconciler[K, T]) onUpdate(ctx context.Context, kind string, newT, registered T) error {
start := time.Now()
err := r.cfg.OnUpdate(ctx, newT, registered)
if err == nil {
r.stats.updated.Add(1)
}
r.observeMetrics(kind, metricLabelOperationUpdate, start, err)
return trace.Wrap(err)
}
// onDelete wraps the OnDelete callback with metrics and stats observation.
func (r *GenericReconciler[K, T]) onDelete(ctx context.Context, kind string, registered T) error {
start := time.Now()
err := r.cfg.OnDelete(ctx, registered)
if err == nil {
r.stats.deleted.Add(1)
}
r.observeMetrics(kind, metricLabelOperationDelete, start, err)
return trace.Wrap(err)
}
func (r *GenericReconciler[K, T]) observeMetrics(kind, operation string, start time.Time, err error) {
r.metrics.reconciliationDuration.With(prometheus.Labels{
metricLabelKind: kind,
metricLabelOperation: operation,
}).Observe(time.Since(start).Seconds())
var result string
switch {
case err == nil:
result = metricLabelResultSuccess
// Only delete-not-found is a noop (resource already gone).
// For create/update, NotFound is a real error (e.g. backend race).
case operation == metricLabelOperationDelete && trace.IsNotFound(err):
result = metricLabelResultNoop
default:
result = metricLabelResultError
}
r.metrics.reconciliationTotal.With(prometheus.Labels{
metricLabelKind: kind,
metricLabelOperation: operation,
metricLabelResult: result,
}).Inc()
}
// Reconcile reconciles currently registered resources with new resources and
// creates/updates/deletes them appropriately.
func (r *GenericReconciler[K, T]) Reconcile(ctx context.Context) error {
r.stats.reset()
currentResources := r.cfg.GetCurrentResources()
newResources := r.cfg.GetNewResources()
r.logger.DebugContext(ctx, "Reconciling current resources with new resources",
"current_resource_count", len(currentResources), "new_resource_count", len(newResources))
start := time.Now()
var errs []error
if r.cfg.Concurrency > 1 {
errs = r.reconcileConcurrent(ctx, currentResources, newResources)
} else {
errs = r.reconcileSequential(ctx, currentResources, newResources)
}
if r.stats.hasChanges() {
r.logger.InfoContext(ctx, "Reconciliation completed",
"kind", r.resourceKind(currentResources, newResources),
"took", time.Since(start).String(),
"stats", r.stats,
)
}
// TODO(zmb3): with a large number of resources, this can return a lengthy
// error message that is difficult to parse
return trace.NewAggregate(errs...)
}
// reconcileSequential processes resources inline without goroutines.
// This is used when Concurrency == 1 (the default) to avoid unnecessary
// goroutine creation overhead.
func (r *GenericReconciler[K, T]) reconcileSequential(ctx context.Context, currentResources, newResources map[K]T) []error {
var errs []error
// Process already registered resources to see if any of them were removed.
for key, current := range currentResources {
if err := r.processRegisteredResource(ctx, newResources, key, current); err != nil {
errs = append(errs, trace.Wrap(err))
}
}
// Add new resources if there are any or refresh those that were updated.
for key, newResource := range newResources {
if err := r.processNewResource(ctx, currentResources, key, newResource); err != nil {
errs = append(errs, trace.Wrap(err))
}
}
return errs
}
// reconcileConcurrent processes resources using an errgroup with the configured
// concurrency limit. This is used when Concurrency > 1.
func (r *GenericReconciler[K, T]) reconcileConcurrent(ctx context.Context, currentResources, newResources map[K]T) []error {
var g errgroup.Group
g.SetLimit(r.cfg.Concurrency)
var (
mu sync.Mutex
errs []error
)
// Process already registered resources to see if any of them were removed.
for key, current := range currentResources {
g.Go(func() error {
if err := r.processRegisteredResource(ctx, newResources, key, current); err != nil {
mu.Lock()
errs = append(errs, trace.Wrap(err))
mu.Unlock()
}
return nil
})
}
// Add new resources if there are any or refresh those that were updated.
for key, newResource := range newResources {
g.Go(func() error {
if err := r.processNewResource(ctx, currentResources, key, newResource); err != nil {
mu.Lock()
errs = append(errs, trace.Wrap(err))
mu.Unlock()
}
return nil
})
}
// Errors are collected separately.
_ = g.Wait()
return errs
}
// resourceKind extracts the resource kind from the first available resource.
func (r *GenericReconciler[K, T]) resourceKind(currentResources, newResources map[K]T) string {
for _, res := range currentResources {
kind, err := types.GetKind(res)
if err == nil {
return kind
}
}
for _, res := range newResources {
kind, err := types.GetKind(res)
if err == nil {
return kind
}
}
return "unknown"
}
// processRegisteredResource checks the specified registered resource against the
// new list of resources.
func (r *GenericReconciler[K, T]) processRegisteredResource(ctx context.Context, newResources map[K]T, key K, registered T) error {
// See if this registered resource is still present among "new" resources.
if _, ok := newResources[key]; ok {
return nil
}
kind, err := types.GetKind(registered)
if err != nil {
return trace.Wrap(err)
}
r.logger.InfoContext(ctx, "Resource was removed, deleting", "kind", kind, "name", key)
err = r.onDelete(ctx, kind, registered)
if err != nil {
if trace.IsNotFound(err) {
r.logger.Log(ctx, logutils.TraceLevel, "Failed to delete resource", "kind", kind, "name", key, "err", err)
return nil
}
return trace.Wrap(err, "failed to delete %v %v", kind, key)
}
return nil
}
// processNewResource checks the provided new resource against currently
// registered resources.
func (r *GenericReconciler[K, T]) processNewResource(ctx context.Context, currentResources map[K]T, key K, newT T) error {
// First see if the resource is already registered and if not, whether it
// matches the selector labels and should be registered.
registered, ok := currentResources[key]
if !ok {
kind, err := types.GetKind(newT)
if err != nil {
return trace.Wrap(err)
}
if r.cfg.Matcher(newT) {
r.logger.InfoContext(ctx, "New resource matches, creating", "kind", kind, "name", key)
if err := r.onCreate(ctx, kind, newT); err != nil {
return trace.Wrap(err, "failed to create %v %v", kind, key)
}
return nil
}
r.logger.DebugContext(ctx, "New resource doesn't match, not creating", "kind", kind, "name", key)
return nil
}
if !r.cfg.AllowOriginChanges {
// Don't overwrite resource of a different origin (e.g., keep static resource from config and ignore dynamic resource)
registeredOrigin, err := types.GetOrigin(registered)
if err != nil {
return trace.Wrap(err)
}
newOrigin, err := types.GetOrigin(newT)
if err != nil {
return trace.Wrap(err)
}
if registeredOrigin != newOrigin {
kind, _ := types.GetKind(newT)
r.logger.WarnContext(ctx, "New resource has different origin, not updating",
"kind", kind, "name", key, "new_origin", newOrigin, "existing_origin", registeredOrigin)
return nil
}
}
// If the resource is already registered but was updated, see if its
// labels still match.
kind, err := types.GetKind(registered)
if err != nil {
return trace.Wrap(err)
}
if r.cfg.CompareResources(newT, registered) != Equal {
if r.cfg.Matcher(newT) {
r.logger.InfoContext(ctx, "Existing resource updated, updating", "kind", kind, "name", key)
if err := r.onUpdate(ctx, kind, newT, registered); err != nil {
return trace.Wrap(err, "failed to update %v %v", kind, key)
}
return nil
}
r.logger.InfoContext(ctx, "Existing resource updated and no longer matches, deleting", "kind", kind, "name", key)
err := r.onDelete(ctx, kind, registered)
if err != nil {
if trace.IsNotFound(err) {
r.logger.Log(ctx, logutils.TraceLevel, "Failed to delete resource", "kind", kind, "name", key, "err", err)
return nil
}
return trace.Wrap(err, "failed to delete %v %v", kind, key)
}
return nil
}
r.logger.Log(ctx, logutils.TraceLevel, "Existing resource is already registered", "kind", kind, "name", key)
return nil
}
// ReconcilerConfig holds the configuration for a reconciler
type ReconcilerConfig[T any] GenericReconcilerConfig[string, T]
// Reconciler reconciles currently registered resources with new resources and
// creates/updates/deletes them appropriately.
//
// This type exists for backwards compatibility, and is a simple wrapper around
// a GenericReconciler[string, T]
type Reconciler[T any] GenericReconciler[string, T]
// NewReconciler creates a new reconciler with provided configuration.
//
// Creates a new GenericReconciler[string, T] and wraps it in a Reconciler[T]
// for backwards compatibility.
func NewReconciler[T any](cfg ReconcilerConfig[T]) (*Reconciler[T], error) {
embedded, err := NewGenericReconciler(GenericReconcilerConfig[string, T](cfg))
if err != nil {
return nil, trace.Wrap(err)
}
return (*Reconciler[T])(embedded), nil
}
// Reconcile reconciles currently registered resources with new resources and
// creates/updates/deletes them appropriately.
func (r *Reconciler[T]) Reconcile(ctx context.Context) error {
return (*GenericReconciler[string, T])(r).Reconcile(ctx)
}
// Teleport
// Copyright (C) 2025 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 (
"github.com/gravitational/trace"
presencev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/presence/v1"
apitypes "github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/common"
)
// ValidateRelayServer will check the given relay server for validity. Should be
// called before writing a new value in the cluster state storage and before
// using a value. The value will not be modified.
func ValidateRelayServer(resource *presencev1.RelayServer) error {
if expected, actual := apitypes.KindRelayServer, resource.GetKind(); expected != actual {
return trace.BadParameter("expected kind %v, got %q", expected, actual)
}
if expected, actual := "", resource.GetSubKind(); expected != actual {
return trace.BadParameter("expected sub_kind %v, got %q", expected, actual)
}
if expected, actual := apitypes.V1, resource.GetVersion(); expected != actual {
return trace.BadParameter("expected version %v, got %q", expected, actual)
}
if name := resource.GetMetadata().GetName(); name == "" {
return trace.BadParameter("missing name")
}
for key := range resource.GetMetadata().GetLabels() {
if key == apitypes.OriginLabel {
return trace.BadParameter("origin label unsupported")
}
if !common.IsValidLabelKey(key) {
return trace.BadParameter("invalid label key %q", key)
}
}
// TODO(espadolini): validate spec contents, nothing to validate so far
return nil
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalRemoteCluster unmarshals the RemoteCluster resource from JSON.
func UnmarshalRemoteCluster(bytes []byte, opts ...MarshalOption) (types.RemoteCluster, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var cluster types.RemoteClusterV3
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
if err := utils.FastUnmarshal(bytes, &cluster); err != nil {
return nil, trace.Wrap(err)
}
err = cluster.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
cluster.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
cluster.SetExpiry(cfg.Expires)
}
return &cluster, nil
}
// MarshalRemoteCluster marshals the RemoteCluster resource to JSON.
func MarshalRemoteCluster(remoteCluster types.RemoteCluster, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch remoteCluster := remoteCluster.(type) {
case *types.RemoteClusterV3:
if err := remoteCluster.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, remoteCluster))
default:
return nil, trace.BadParameter("unrecognized remote cluster version %T", remoteCluster)
}
}
/*
* 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 (
"encoding/json"
"fmt"
"strings"
"sync"
"time"
"github.com/gravitational/trace"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/protoadapt"
"google.golang.org/protobuf/types/known/timestamppb"
autoupdatev1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/autoupdate/v1"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
healthcheckconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/healthcheckconfig/v1"
machineidv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
subcav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/subca/v1"
workloadidentityv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
scopedaccess "github.com/gravitational/teleport/lib/scopes/access"
"github.com/gravitational/teleport/lib/utils"
)
// MarshalConfig specifies marshaling options
type MarshalConfig struct {
// Version specifies a particular version we should marshal resources with
Version string
// Revision of the resource to assign.
Revision string
// PreserveRevision preserves revision in resource
// specs when marshaling
PreserveRevision bool
// Expires is an optional expiry time
Expires time.Time
// DisallowUnknown will, for resources stored in protojson, disallow unknown
// fields when unmarshaling. This is useful if a resource is being parsed
// from user-specified data rather than persistent cluster state storage.
DisallowUnknown bool
}
// GetVersion returns explicitly provided version or sets latest as default
func (m *MarshalConfig) GetVersion() string {
if m.Version == "" {
return types.V2
}
return m.Version
}
// MarshalOption sets marshaling option
type MarshalOption func(c *MarshalConfig) error
// CollectOptions collects all options from functional arg and returns config
func CollectOptions(opts []MarshalOption) (*MarshalConfig, error) {
var cfg MarshalConfig
for _, o := range opts {
if err := o(&cfg); err != nil {
return nil, trace.Wrap(err)
}
}
return &cfg, nil
}
// AddOptions adds marshal options and returns a new copy
func AddOptions(opts []MarshalOption, add ...MarshalOption) []MarshalOption {
out := make([]MarshalOption, len(opts), len(opts)+len(add))
copy(out, opts)
return append(opts, add...)
}
// WithRevision assigns Revision to the resource
func WithRevision(rev string) MarshalOption {
return func(c *MarshalConfig) error {
c.Revision = rev
return nil
}
}
// WithExpires assigns expiry value
func WithExpires(expires time.Time) MarshalOption {
return func(c *MarshalConfig) error {
c.Expires = expires
return nil
}
}
// WithVersion sets marshal version
func WithVersion(v string) MarshalOption {
return func(c *MarshalConfig) error {
switch v {
case types.V1, types.V2, types.V3:
c.Version = v
return nil
default:
return trace.BadParameter("version '%v' is not supported", v)
}
}
}
// DisallowUnknown will, for resources stored in protojson, disallow unknown
// fields when unmarshaling. This is useful if a resource is being parsed
// from user-specified data rather than persistent cluster state storage.
func DisallowUnknown() MarshalOption {
return func(c *MarshalConfig) error {
c.DisallowUnknown = true
return nil
}
}
// PreserveRevision preserves revision when
// marshaling value
func PreserveRevision() MarshalOption {
return func(c *MarshalConfig) error {
c.PreserveRevision = true
return nil
}
}
// ParseShortcut parses resource shortcut
// Generally, this should include the plural of a singular resource name or vice
// versa.
func ParseShortcut(in string) (string, error) {
if in == "" {
return "", trace.BadParameter("missing resource name")
}
switch strings.ToLower(in) {
case types.KindRole, "roles":
return types.KindRole, nil
case types.KindNamespace, "namespaces", "ns":
return types.KindNamespace, nil
case types.KindAuthServer, "auth_servers", "auth":
return types.KindAuthServer, nil
case types.KindProxy, "proxies":
return types.KindProxy, nil
case types.KindNode, "nodes":
return types.KindNode, nil
case types.KindOIDCConnector:
return types.KindOIDCConnector, nil
case types.KindSAMLConnector:
return types.KindSAMLConnector, nil
case types.KindGithubConnector:
return types.KindGithubConnector, nil
case types.KindConnectors, "connector":
return types.KindConnectors, nil
case types.KindUser, "users":
return types.KindUser, nil
case types.KindCertAuthority, "cert_authorities", "cas":
return types.KindCertAuthority, nil
case types.KindCertAuthorityOverride, "cert_authority_overrides", "ca_override", "ca_overrides":
return types.KindCertAuthorityOverride, nil
case types.KindClientIPRestriction, types.KindClientIPRestriction + "s":
return types.KindClientIPRestriction, nil
case types.KindReverseTunnel, "reverse_tunnels", "rts":
return types.KindReverseTunnel, nil
case types.KindTrustedCluster, "tc", "cluster", "clusters":
return types.KindTrustedCluster, nil
case types.KindClusterAuthPreference, "cluster_authentication_preferences", "cluster_auth_preferences", "cap":
return types.KindClusterAuthPreference, nil
case types.KindUIConfig, "ui":
return types.KindUIConfig, nil
case types.KindClusterNetworkingConfig, "networking_config", "networking", "net_config", "netconfig":
return types.KindClusterNetworkingConfig, nil
case types.KindSessionRecordingConfig, "recording_config", "session_recording", "rec_config", "recconfig":
return types.KindSessionRecordingConfig, nil
case types.KindExternalAuditStorage:
return types.KindExternalAuditStorage, nil
case types.KindRemoteCluster, "remote_clusters", "rc", "rcs":
return types.KindRemoteCluster, nil
case types.KindSemaphore, "semaphores", "sem", "sems":
return types.KindSemaphore, nil
case types.KindKubernetesCluster, "kube_clusters":
return types.KindKubernetesCluster, nil
case types.KindKubeServer, "kube_servers":
return types.KindKubeServer, nil
case types.KindLock, "locks":
return types.KindLock, nil
case types.KindDatabaseServer, "db_servers":
return types.KindDatabaseServer, nil
case types.KindNetworkRestrictions:
return types.KindNetworkRestrictions, nil
case types.KindDatabase:
return types.KindDatabase, nil
case types.KindApp, "apps":
return types.KindApp, nil
case types.KindAppServer, "app_servers":
return types.KindAppServer, nil
case types.KindWindowsDesktopService, "windows_service", "win_desktop_service", "win_service", "windows_desktop_services":
return types.KindWindowsDesktopService, nil
case types.KindWindowsDesktop, "win_desktop":
return types.KindWindowsDesktop, nil
case types.KindDynamicWindowsDesktop, "dynamic_win_desktop", "dynamic_desktop":
return types.KindDynamicWindowsDesktop, nil
case types.KindLinuxDesktop, types.KindLinuxDesktop + "s":
return types.KindLinuxDesktop, nil
case types.KindToken, "tokens":
return types.KindToken, nil
case types.KindInstaller:
return types.KindInstaller, nil
case types.KindDatabaseService, types.KindDatabaseService + "s":
return types.KindDatabaseService, nil
case types.KindLoginRule, types.KindLoginRule + "s":
return types.KindLoginRule, nil
case types.KindSAMLIdPServiceProvider, types.KindSAMLIdPServiceProvider + "s", "saml_sp", "saml_sps":
return types.KindSAMLIdPServiceProvider, nil
case types.KindUserGroup, types.KindUserGroup + "s", "usergroup", "usergroups":
return types.KindUserGroup, nil
case types.KindDevice, types.KindDevice + "s":
return types.KindDevice, nil
case types.KindOktaImportRule, types.KindOktaImportRule + "s", "oktaimportrule", "oktaimportrules":
return types.KindOktaImportRule, nil
case types.KindOktaAssignment, types.KindOktaAssignment + "s", "oktaassignment", "oktaassignments":
return types.KindOktaAssignment, nil
case types.KindClusterMaintenanceConfig, "cmc":
return types.KindClusterMaintenanceConfig, nil
case types.KindIntegration, types.KindIntegration + "s":
return types.KindIntegration, nil
case types.KindAccessList, types.KindAccessList + "s", "accesslist", "accesslists":
return types.KindAccessList, nil
case types.KindDiscoveryConfig, types.KindDiscoveryConfig + "s", "discoveryconfig", "discoveryconfigs":
return types.KindDiscoveryConfig, nil
case types.KindAuditQuery:
return types.KindAuditQuery, nil
case types.KindSecurityReport:
return types.KindSecurityReport, nil
case types.KindServerInfo:
return types.KindServerInfo, nil
case types.KindBot, "bots":
return types.KindBot, nil
case types.KindBotInstance, types.KindBotInstance + "s":
return types.KindBotInstance, nil
case types.KindDatabaseObjectImportRule, "db_object_import_rules", "database_object_import_rule":
return types.KindDatabaseObjectImportRule, nil
case types.KindAccessMonitoringRule:
return types.KindAccessMonitoringRule, nil
case types.KindDatabaseObject, "database_object":
return types.KindDatabaseObject, nil
case types.KindCrownJewel, "crown_jewels":
return types.KindCrownJewel, nil
case types.KindVnetConfig:
return types.KindVnetConfig, nil
case types.KindAccessRequest, types.KindAccessRequest + "s", "accessrequest", "accessrequests":
return types.KindAccessRequest, nil
case types.KindPlugin, types.KindPlugin + "s":
return types.KindPlugin, nil
case types.KindAccessGraphSettings, "ags":
return types.KindAccessGraphSettings, nil
case types.KindSPIFFEFederation, types.KindSPIFFEFederation + "s":
return types.KindSPIFFEFederation, nil
case types.KindWorkloadIdentity, types.KindWorkloadIdentity + "s", "workload_identities", "workloadidentity", "workloadidentities", "workloadidentitys":
return types.KindWorkloadIdentity, nil
case types.KindStaticHostUser, types.KindStaticHostUser + "s", "host_user", "host_users":
return types.KindStaticHostUser, nil
case types.KindUserTask, types.KindUserTask + "s":
return types.KindUserTask, nil
case types.KindAutoUpdateConfig:
return types.KindAutoUpdateConfig, nil
case types.KindAutoUpdateVersion:
return types.KindAutoUpdateVersion, nil
case types.KindAutoUpdateAgentRollout:
return types.KindAutoUpdateAgentRollout, nil
case types.KindAutoUpdateAgentReport:
return types.KindAutoUpdateAgentReport, nil
case types.KindAutoUpdateBotInstanceReport:
return types.KindAutoUpdateBotInstanceReport, nil
case types.KindGitServer, types.KindGitServer + "s":
return types.KindGitServer, nil
case types.KindWorkloadIdentityX509Revocation, types.KindWorkloadIdentityX509Revocation + "s":
return types.KindWorkloadIdentityX509Revocation, nil
case types.KindWorkloadIdentityX509IssuerOverride, types.KindWorkloadIdentityX509IssuerOverride + "s":
return types.KindWorkloadIdentityX509IssuerOverride, nil
case types.KindSigstorePolicy, "sigstorepolicy", "sigstore_policies", "sigstorepolicies":
return types.KindSigstorePolicy, nil
case types.KindHealthCheckConfig, types.KindHealthCheckConfig + "s", "hcc":
return types.KindHealthCheckConfig, nil
case scopedaccess.KindScopedRole, scopedaccess.KindScopedRole + "s", "scopedrole", "scopedroles":
return scopedaccess.KindScopedRole, nil
case scopedaccess.KindScopedRoleAssignment, scopedaccess.KindScopedRoleAssignment + "s", "scopedroleassignment", "scopedroleassignments", "sra":
return scopedaccess.KindScopedRoleAssignment, nil
case types.KindInferenceModel, "inference_models":
return types.KindInferenceModel, nil
case types.KindInferenceSecret, "inference_secrets":
return types.KindInferenceSecret, nil
case types.KindInferencePolicy, "inference_policies":
return types.KindInferencePolicy, nil
case types.KindClassifier, types.KindClassifier + "s":
return types.KindClassifier, nil
case types.KindRetrievalModel:
return types.KindRetrievalModel, nil
case types.KindRelayServer, types.KindRelayServer + "s":
return types.KindRelayServer, nil
case types.KindWorkloadCluster, types.KindWorkloadCluster + "s":
return types.KindWorkloadCluster, nil
case scopedaccess.KindScopedToken, scopedaccess.KindScopedToken + "s", "scopedtoken", "scopedtokens":
return scopedaccess.KindScopedToken, nil
case types.KindBeamsConfig:
return types.KindBeamsConfig, nil
}
return "", trace.BadParameter("unsupported resource: %q - resources should be expressed as 'type/name', for example 'connector/github'", in)
}
// ParseRef parses resource reference eg daemonsets/ds1
func ParseRef(ref string) (*Ref, error) {
if ref == "" {
return nil, trace.BadParameter("missing value")
}
parts := strings.FieldsFunc(ref, isDelimiter)
switch len(parts) {
case 1:
shortcut, err := ParseShortcut(parts[0])
if err != nil {
return nil, trace.Wrap(err)
}
return &Ref{Kind: shortcut}, nil
case 2:
shortcut, err := ParseShortcut(parts[0])
if err != nil {
return nil, trace.Wrap(err)
}
return &Ref{Kind: shortcut, Name: parts[1]}, nil
case 3:
shortcut, err := ParseShortcut(parts[0])
if err != nil {
return nil, trace.Wrap(err)
}
return &Ref{Kind: shortcut, SubKind: parts[1], Name: parts[2]}, nil
}
return nil, trace.BadParameter("failed to parse '%v'", ref)
}
// isDelimiter returns true if rune is space or /
func isDelimiter(r rune) bool {
switch r {
case '\t', ' ', '/':
return true
}
return false
}
// Ref is a resource reference. Typically of the form kind/name,
// but sometimes of the form kind/subkind/name.
type Ref struct {
Kind string
SubKind string
Name string
}
// Set sets the name of the resource
func (r *Ref) Set(v string) error {
out, err := ParseRef(v)
if err != nil {
return err
}
*r = *out
return nil
}
func (r *Ref) String() string {
if r.SubKind == "" {
if r.Name == "" {
return r.Kind
}
return fmt.Sprintf("%s/%s", r.Kind, r.Name)
}
return fmt.Sprintf("%s/%s/%s", r.Kind, r.SubKind, r.Name)
}
// Refs is a set of resource references
type Refs []Ref
// ParseRefs parses a comma-separated string of resource references (eg "users/alice,users/bob")
func ParseRefs(refs string) (Refs, error) {
if refs == "all" {
return []Ref{{Kind: "all"}}, nil
}
var escaped bool
isBreak := func(r rune) bool {
brk := false
switch r {
case ',':
brk = true && !escaped
escaped = false
case '\\':
escaped = true && !escaped
default:
escaped = false
}
return brk
}
var parsed []Ref
split := fieldsFunc(strings.TrimSpace(refs), isBreak)
for _, s := range split {
ref, err := ParseRef(strings.ReplaceAll(s, `\,`, `,`))
if err != nil {
return nil, trace.Wrap(err)
}
parsed = append(parsed, *ref)
}
return parsed, nil
}
// Set sets the value of `r` from a comma-separated string of resource
// references (in-place equivalent of `ParseRefs`).
func (r *Refs) Set(v string) error {
refs, err := ParseRefs(v)
if err != nil {
return trace.Wrap(err)
}
*r = refs
return nil
}
// IsAll checks if refs is special wildcard case `all`.
func (r *Refs) IsAll() bool {
refs := *r
if len(refs) != 1 {
return false
}
return refs[0].Kind == "all"
}
func (r *Refs) String() string {
var builder strings.Builder
for i, ref := range *r {
if i > 0 {
builder.WriteRune(',')
}
builder.WriteString(ref.String())
}
return builder.String()
}
// fieldsFunc is an exact copy of the current implementation of `strings.FieldsFunc`.
// The docs of `strings.FieldsFunc` indicate that future implementations may not call
// `f` on every rune, may not preserve ordering, or may panic if `f` does not return the
// same output for every instance of a given rune. All of these changes would break
// our implementation of backslash-escaping, so we're using a local copy.
func fieldsFunc(s string, f func(rune) bool) []string {
// A span is used to record a slice of s of the form s[start:end].
// The start index is inclusive and the end index is exclusive.
type span struct {
start int
end int
}
spans := make([]span, 0, 32)
// Find the field start and end indices.
wasField := false
fromIndex := 0
for i, rune := range s {
if f(rune) {
if wasField {
spans = append(spans, span{start: fromIndex, end: i})
wasField = false
}
} else {
if !wasField {
fromIndex = i
wasField = true
}
}
}
// Last field might end at EOF.
if wasField {
spans = append(spans, span{fromIndex, len(s)})
}
// Create strings from recorded field indices.
a := make([]string, len(spans))
for i, span := range spans {
a[i] = s[span.start:span.end]
}
return a
}
// marshalerMutex is a mutex for resource marshalers/unmarshalers
var marshalerMutex sync.RWMutex
// ResourceMarshaler handles marshaling of a specific resource type.
type ResourceMarshaler func(types.Resource, ...MarshalOption) ([]byte, error)
// ResourceUnmarshaler handles unmarshaling of a specific resource type.
type ResourceUnmarshaler func([]byte, ...MarshalOption) (types.Resource, error)
// resourceMarshalers holds a collection of marshalers organized by kind.
var resourceMarshalers = make(map[string]ResourceMarshaler)
// resourceUnmarshalers holds a collection of unmarshalers organized by kind.
var resourceUnmarshalers = make(map[string]ResourceUnmarshaler)
// GetResourceMarshalerKinds lists all registered resource marshalers by kind.
func GetResourceMarshalerKinds() []string {
marshalerMutex.Lock()
defer marshalerMutex.Unlock()
kinds := make([]string, 0, len(resourceMarshalers))
for kind := range resourceMarshalers {
kinds = append(kinds, kind)
}
return kinds
}
// RegisterResourceMarshaler registers a marshaler for resources of a specific kind.
// WARNING!!
// Registering a resource Marshaler requires lib/services/local.CreateResources
// supports the resource kind or the standard backup/restore procedure of using
// `tctl get all` and then BootstrapResources in Teleport will fail.
func RegisterResourceMarshaler(kind string, marshaler ResourceMarshaler) {
marshalerMutex.Lock()
defer marshalerMutex.Unlock()
resourceMarshalers[kind] = marshaler
}
// RegisterResourceUnmarshaler registers an unmarshaler for resources of a specific kind.
func RegisterResourceUnmarshaler(kind string, unmarshaler ResourceUnmarshaler) {
marshalerMutex.Lock()
defer marshalerMutex.Unlock()
resourceUnmarshalers[kind] = unmarshaler
}
func getResourceMarshaler(kind string) (ResourceMarshaler, bool) {
marshalerMutex.RLock()
defer marshalerMutex.RUnlock()
m, ok := resourceMarshalers[kind]
if !ok {
return nil, false
}
return m, true
}
func getResourceUnmarshaler(kind string) (ResourceUnmarshaler, bool) {
marshalerMutex.RLock()
defer marshalerMutex.RUnlock()
u, ok := resourceUnmarshalers[kind]
if !ok {
return nil, false
}
return u, true
}
func init() {
RegisterResourceMarshaler(types.KindUser, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
user, ok := resource.(types.User)
if !ok {
return nil, trace.BadParameter("expected User, got %T", resource)
}
bytes, err := MarshalUser(user, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindUser, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
user, err := UnmarshalUser(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return user, nil
})
RegisterResourceMarshaler(types.KindCertAuthority, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
certAuthority, ok := resource.(types.CertAuthority)
if !ok {
return nil, trace.BadParameter("expected CertAuthority, got %T", resource)
}
bytes, err := MarshalCertAuthority(certAuthority, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindCertAuthority, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
certAuthority, err := UnmarshalCertAuthority(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return certAuthority, nil
})
RegisterResourceMarshaler(types.KindTrustedCluster, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
trustedCluster, ok := resource.(types.TrustedCluster)
if !ok {
return nil, trace.BadParameter("expected TrustedCluster, got %T", resource)
}
bytes, err := MarshalTrustedCluster(trustedCluster, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindTrustedCluster, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
trustedCluster, err := UnmarshalTrustedCluster(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return trustedCluster, nil
})
RegisterResourceMarshaler(types.KindGithubConnector, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
githubConnector, ok := resource.(types.GithubConnector)
if !ok {
return nil, trace.BadParameter("expected GithubConnector, got %T", resource)
}
bytes, err := MarshalOSSGithubConnector(githubConnector, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindGithubConnector, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
githubConnector, err := UnmarshalOSSGithubConnector(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return githubConnector, nil
})
RegisterResourceMarshaler(types.KindSAMLConnector, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
samlConnector, ok := resource.(types.SAMLConnector)
if !ok {
return nil, trace.BadParameter("expected SAMLConnector, got %T", resource)
}
bytes, err := MarshalSAMLConnector(samlConnector, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindSAMLConnector, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
samlConnector, err := UnmarshalSAMLConnector(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return samlConnector, nil
})
RegisterResourceMarshaler(types.KindOIDCConnector, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
oidConnector, ok := resource.(types.OIDCConnector)
if !ok {
return nil, trace.BadParameter("expected OIDCConnector, got %T", resource)
}
bytes, err := MarshalOIDCConnector(oidConnector, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindOIDCConnector, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
oidcConnector, err := UnmarshalOIDCConnector(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return oidcConnector, nil
})
RegisterResourceMarshaler(types.KindRole, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
role, ok := resource.(types.Role)
if !ok {
return nil, trace.BadParameter("expected Role, got %T", resource)
}
bytes, err := MarshalRole(role, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindRole, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
role, err := UnmarshalRole(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return role, nil
})
RegisterResourceMarshaler(types.KindToken, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
token, ok := resource.(types.ProvisionToken)
if !ok {
return nil, trace.BadParameter("expected Token, got %T", resource)
}
bytes, err := MarshalProvisionToken(token, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindToken, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
token, err := UnmarshalProvisionToken(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return token, nil
})
RegisterResourceMarshaler(types.KindLock, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
lock, ok := resource.(types.Lock)
if !ok {
return nil, trace.BadParameter("expected lock, got %T", resource)
}
bytes, err := MarshalLock(lock, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindLock, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
lock, err := UnmarshalLock(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return lock, nil
})
RegisterResourceMarshaler(types.KindClusterNetworkingConfig, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
cnc, ok := resource.(types.ClusterNetworkingConfig)
if !ok {
return nil, trace.BadParameter("expected cluster_networking_config go %T", resource)
}
bytes, err := MarshalClusterNetworkingConfig(cnc, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindClusterNetworkingConfig, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
cnc, err := UnmarshalClusterNetworkingConfig(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return cnc, nil
})
RegisterResourceMarshaler(types.KindClusterAuthPreference, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
ap, ok := resource.(types.AuthPreference)
if !ok {
return nil, trace.BadParameter("expected cluster_auth_preference go %T", resource)
}
bytes, err := MarshalAuthPreference(ap, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindClusterAuthPreference, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
ap, err := UnmarshalAuthPreference(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return ap, nil
})
RegisterResourceUnmarshaler(types.KindBot, func(bytes []byte, option ...MarshalOption) (types.Resource, error) {
cfg, err := CollectOptions(option)
if err != nil {
return nil, err
}
b := &machineidv1pb.Bot{}
if err := (protojson.UnmarshalOptions{DiscardUnknown: !cfg.DisallowUnknown}).Unmarshal(bytes, b); err != nil {
return nil, trace.Wrap(err)
}
return types.Resource153ToLegacy(b), nil
})
RegisterResourceUnmarshaler(types.KindAutoUpdateConfig, func(bytes []byte, option ...MarshalOption) (types.Resource, error) {
cfg, err := CollectOptions(option)
if err != nil {
return nil, err
}
c := &autoupdatev1pb.AutoUpdateConfig{}
if err := (protojson.UnmarshalOptions{DiscardUnknown: !cfg.DisallowUnknown}).Unmarshal(bytes, c); err != nil {
return nil, trace.Wrap(err)
}
return types.Resource153ToLegacy(c), nil
})
RegisterResourceUnmarshaler(types.KindAutoUpdateVersion, func(bytes []byte, option ...MarshalOption) (types.Resource, error) {
cfg, err := CollectOptions(option)
if err != nil {
return nil, err
}
v := &autoupdatev1pb.AutoUpdateVersion{}
if err := (protojson.UnmarshalOptions{DiscardUnknown: !cfg.DisallowUnknown}).Unmarshal(bytes, v); err != nil {
return nil, trace.Wrap(err)
}
return types.Resource153ToLegacy(v), nil
})
// add health_check_config to tctl get all
RegisterResourceMarshaler(types.KindHealthCheckConfig, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
wrapper, ok := resource.(types.Resource153UnwrapperT[*healthcheckconfigv1.HealthCheckConfig])
if !ok {
return nil, trace.BadParameter("expected health check config, got %T", resource)
}
bytes, err := MarshalHealthCheckConfig(wrapper.UnwrapT(), opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
// support health_check_config --bootstrap and --apply-on-startup
RegisterResourceUnmarshaler(types.KindHealthCheckConfig, func(data []byte, options ...MarshalOption) (types.Resource, error) {
cfg, err := UnmarshalHealthCheckConfig(data, options...)
if err != nil {
return nil, trace.Wrap(err)
}
return types.Resource153ToLegacy(cfg), nil
})
RegisterResourceUnmarshaler(types.KindWorkloadIdentity, func(bytes []byte, option ...MarshalOption) (types.Resource, error) {
cfg, err := CollectOptions(option)
if err != nil {
return nil, err
}
wid := &workloadidentityv1.WorkloadIdentity{}
if err := (protojson.UnmarshalOptions{DiscardUnknown: !cfg.DisallowUnknown}).Unmarshal(bytes, wid); err != nil {
return nil, trace.Wrap(err)
}
return types.Resource153ToLegacy(wid), nil
})
RegisterResourceMarshaler(types.KindCertAuthorityOverride, func(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
unwrapper, ok := resource.(types.Resource153UnwrapperT[*subcav1.CertAuthorityOverride])
if !ok {
return nil, trace.BadParameter("expected wrapped CertAuthorityOverride resource, got %T", resource)
}
caOverride := unwrapper.UnwrapT()
if caOverride == nil {
return nil, trace.BadParameter("nil CertAuthorityOverride resource")
}
bytes, err := MarshalCertAuthorityOverride(caOverride, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
})
RegisterResourceUnmarshaler(types.KindCertAuthorityOverride, func(bytes []byte, opts ...MarshalOption) (types.Resource, error) {
caOverride, err := UnmarshalCertAuthorityOverride(bytes, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return types.ProtoResource153ToLegacy(caOverride), nil
})
}
// CheckAndSetDefaults calls [r.CheckAndSetDefaults] if r implements the method.
// If r does not implement, then this is a nop.
//
// This method exists for backwards compatibility with old-style resources.
// Prefer using RFD 153 style resources, passing concrete types and running
// validations before storage writes only.
func CheckAndSetDefaults(r any) error {
if r, ok := r.(interface{ CheckAndSetDefaults() error }); ok {
return trace.Wrap(r.CheckAndSetDefaults())
}
return nil
}
// MarshalResource attempts to marshal a resource dynamically, returning NotImplementedError
// if no marshaler has been registered.
//
// NOTE: This function only supports the subset of resources which may be imported/exported
// by users (e.g. via `tctl get`).
func MarshalResource(resource types.Resource, opts ...MarshalOption) ([]byte, error) {
marshal, ok := getResourceMarshaler(resource.GetKind())
if !ok {
return nil, trace.NotImplemented("cannot dynamically marshal resources of kind %q", resource.GetKind())
}
// Handle the case where `resource` was never fully unmarshaled.
if r, ok := resource.(*UnknownResource); ok {
u, err := UnmarshalResource(r.GetKind(), r.Raw, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
resource = u
}
m, err := marshal(resource, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return m, nil
}
// UnmarshalResource attempts to unmarshal a resource dynamically, returning NotImplementedError
// if no unmarshaler has been registered.
//
// NOTE: This function only supports the subset of resources which may be imported/exported
// by users (e.g. via `tctl get`).
func UnmarshalResource(kind string, raw []byte, opts ...MarshalOption) (types.Resource, error) {
unmarshal, ok := getResourceUnmarshaler(kind)
if !ok {
return nil, trace.NotImplemented("cannot dynamically unmarshal resources of kind %q", kind)
}
u, err := unmarshal(raw, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
return u, nil
}
// UnknownResource is used to detect resources
type UnknownResource struct {
types.ResourceHeader
// Raw is raw representation of the resource
Raw []byte
}
// UnmarshalJSON unmarshals header and captures raw state
func (u *UnknownResource) UnmarshalJSON(raw []byte) error {
var h types.ResourceHeader
if err := json.Unmarshal(raw, &h); err != nil {
return trace.Wrap(err)
}
u.Raw = make([]byte, len(raw))
u.ResourceHeader = h
copy(u.Raw, raw)
return nil
}
// setResourceName modifies the types.Metadata argument in place, setting the resource name.
// The name is calculated based on nameParts arguments which are joined by hyphens "-".
// If a name override label is present, it will replace the *first* name part.
func setResourceName(overrideLabels []string, meta types.Metadata, firstNamePart string, extraNameParts ...string) types.Metadata {
nameParts := append([]string{firstNamePart}, extraNameParts...)
// apply override
for _, overrideLabel := range overrideLabels {
if override, found := meta.Labels[overrideLabel]; found && override != "" {
nameParts[0] = override
break
}
}
meta.Name = strings.Join(nameParts, "-")
return meta
}
type resetProtoResource interface {
protoadapt.MessageV1
SetRevision(string)
}
// maybeResetProtoRevision returns a clone of [r] with the identifiers reset to default values if
// preserveRevision is true, otherwise this is a nop, and the original value is returned unaltered.
//
// Prefer maybeResetProtoRevisionv2 for newer RFD153-style resources, only one or the other should compile
// for any given type.
func maybeResetProtoRevision[T resetProtoResource](preserveRevision bool, r T) T {
if preserveRevision {
return r
}
cp := apiutils.CloneProtoMsg(r)
cp.SetRevision("")
return cp
}
// ProtoResource abstracts a resource defined as a protobuf message.
type ProtoResource interface {
proto.Message
// GetMetadata returns the generic resource metadata.
GetMetadata() *headerv1.Metadata
}
// ProtoResourcePtr is a ProtoResource that is also a pointer to T.
type ProtoResourcePtr[T any] interface {
*T
ProtoResource
}
// maybeResetProtoRevisionv2 returns a clone of [r] with the identifiers reset to default values if
// preserveRevision is true, otherwise this is a nop, and the original value is returned unaltered.
//
// This is like maybeResetProtoRevision but made for newer RFD153-style resources which implement a
// different interface, only one or the other should compile for any given type.
func maybeResetProtoRevisionv2[T ProtoResource](preserveRevision bool, r T) T {
if preserveRevision {
return r
}
cp := proto.Clone(r).(T)
cp.GetMetadata().SetRevision("")
return cp
}
// MarshalProtoResource marshals a ProtoResource to JSON using [protojson.Marshal] and respecting [opts].
func MarshalProtoResource[T ProtoResource](resource T, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
resource = maybeResetProtoRevisionv2(cfg.PreserveRevision, resource)
data, err := protojson.Marshal(resource)
if err != nil {
return nil, trace.Wrap(err)
}
return data, nil
}
// UnmarshalProtoResource unmarshals a ProtoResource from JSON using [protojson.Unmarshal] and respecting [opts].
// It is paramaterized on types T and U, where T is a pointer type that implements ProtoResource, and U is the
// type that T points to. This is so that it can allocate an instance of U to unmarshal into without
// reflection.
func UnmarshalProtoResource[T ProtoResourcePtr[U], U any](data []byte, opts ...MarshalOption) (T, error) {
if len(data) == 0 {
return nil, trace.BadParameter("nothing to unmarshal")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var resource T = new(U)
err = protojson.UnmarshalOptions{DiscardUnknown: !cfg.DisallowUnknown}.Unmarshal(data, resource)
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
resource.GetMetadata().SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
resource.GetMetadata().SetExpires(timestamppb.New(cfg.Expires))
}
return resource, nil
}
// UnmarshalProtoResourceArray unmarshals an array of ProtoResources from JSON using [UnmarshalProtoResource] on each
// individual element.
func UnmarshalProtoResourceArray[T ProtoResourcePtr[U], U any](data []byte, opts ...MarshalOption) ([]T, error) {
var msgs []json.RawMessage
if err := json.Unmarshal(data, &msgs); err != nil {
return nil, trace.Wrap(err)
}
resources := make([]T, 0, len(msgs))
for _, msg := range msgs {
resource, err := UnmarshalProtoResource[T](msg, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
resources = append(resources, resource)
}
return resources, nil
}
// FastMarshalProtoResourceDeprecated marshals a ProtoResource to JSON using [utils.FastMarshal] and respecting [opts].
//
// Deprecated: this should not be used for new types, prefer [MarshalProtoResource]. Existing types should not
// be converted to maintain compatibility.
func FastMarshalProtoResourceDeprecated[T ProtoResource](resource T, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
resource = maybeResetProtoRevisionv2(cfg.PreserveRevision, resource)
data, err := utils.FastMarshal(resource)
if err != nil {
return nil, trace.Wrap(err)
}
return data, nil
}
// FastUnmarshalProtoResourceDeprecated unmarshals a ProtoResource from JSON using [utils.FastUnmarshal] and respecting [opts].
// It is paramaterized on types T and U, where T is a pointer type that implements ProtoResource, and U is the
// type that T points to. This is so that it can allocate an instance of U to unmarshal into without
// reflection.
//
// Deprecated: this should not be used for new types, prefer [UnmarshalProtoResource]. Existing types should not
// be converted to maintain compatibility.
func FastUnmarshalProtoResourceDeprecated[T ProtoResourcePtr[U], U any](data []byte, opts ...MarshalOption) (T, error) {
if len(data) == 0 {
return nil, trace.BadParameter("nothing to unmarshal")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var resource T = new(U)
err = utils.FastUnmarshal(data, resource)
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
resource.GetMetadata().SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
resource.GetMetadata().SetExpires(timestamppb.New(cfg.Expires))
}
return resource, nil
}
/*
* Teleport
* Copyright (C) 2025 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils/set"
)
// MatcherTransform defines a func wrapping a RoleMatcher to modify or extend its behavior.
type MatcherTransform func(RoleMatcher) RoleMatcher
// WithConstraints returns a MatcherTransform that scopes principal-bearing
// RoleMatchers to any provided ResourceConstraints.
//
// For matchers that encode a specific principal (e.g., AWS Role ARN, IC assignment,
// SSH login), the returned transform first checks that principal against the provided
// ResourceConstraints; if it's not present, the transformed matcher fails fast. If it is
// present, the original matcher's logic is applied.
//
// For non-principal-bearing matchers, the transform is a no-op.
//
// This enforces that even if a role would otherwise match a principal on a
// resource, the principal must also be allowed by the resource's Constraints.
func WithConstraints(rc *types.ResourceConstraints) MatcherTransform {
if rc == nil {
return func(m RoleMatcher) RoleMatcher { return m }
}
switch d := rc.Details.(type) {
case *types.ResourceConstraints_AwsConsole:
return buildStringConstraintTransform(
d.Validate,
func() []string { return d.AwsConsole.RoleArns },
func(m RoleMatcher) string {
principal := ""
switch lm := m.(type) {
case *awsAppLoginMatcher:
principal = lm.awsRole
case *AWSRoleARNMatcher:
principal = lm.RoleARN
}
return principal
},
)
case *types.ResourceConstraints_Ssh:
return buildStringConstraintTransform(
d.Validate,
func() []string { return d.Ssh.Logins },
func(m RoleMatcher) string {
lm, ok := m.(*loginMatcher)
if !ok {
return ""
}
return lm.login
},
)
// TODO(kiosion): Future support for AWS Identity Center.
// Need to decide on best way to handle; whether to continue using IdentityCenterAccountAssignments, or just Account, with PermissionSets carried in constraints.
default:
return func(m RoleMatcher) RoleMatcher {
return RoleMatcherFunc(func(_ types.Role, _ types.RoleConditionType) (bool, error) {
return false, trace.BadParameter("unsupported constraint details type %T", d)
})
}
}
}
// buildStringConstraintTransform factors out shared logic for string-list-based
// ResourceConstraints (e.g., AWS role ARNs, SSH logins). It handles validation,
// then builds the principal-gated RoleMatcher transform.
func buildStringConstraintTransform(
validate func() error,
getStrings func() []string,
getPrincipal func(RoleMatcher) string,
) MatcherTransform {
if err := validate(); err != nil {
return func(m RoleMatcher) RoleMatcher {
return RoleMatcherFunc(func(_ types.Role, _ types.RoleConditionType) (bool, error) {
return false, trace.Wrap(err)
})
}
}
allowedSet := set.New(getStrings()...)
return func(m RoleMatcher) RoleMatcher {
principal := getPrincipal(m)
if principal == "" {
return m // non-principal-bearing matcher; no-op
}
return RoleMatcherFunc(func(role types.Role, cond types.RoleConditionType) (bool, error) {
if !allowedSet.Contains(principal) {
return false, nil
}
return m.Match(role, cond)
})
}
}
// BuildResourceConstraintMatchers returns RoleMatchers derived from any
// ResourceConstraints requested for the given resource, correlating the
// resource against resourceAccessIDs by kind and name. Entries without
// constraints contribute no matchers, so resource kinds that cannot carry
// constraints are unaffected.
//
// Correlating by kind and name mirrors how requested resources are looked up
// from their IDs (see [accessrequest.GetResourcesByResourceIDs]); callers are
// expected to pass resources and resourceAccessIDs scoped to the same cluster.
//
// TODO(kiosion): When constraints extend for Kubernetes support, kube sub-resource
// IDs need name-only correlation against the kube_cluster resource, like
// getKubeResourcesFromResourceIDs
func BuildResourceConstraintMatchers(resourceAccessIDs []types.ResourceAccessID, resource types.Resource) ([]RoleMatcher, error) {
var matchers []RoleMatcher
for _, raid := range resourceAccessIDs {
rid := raid.GetResourceID()
if rid.Name != resource.GetName() || rid.Kind != resource.GetKind() {
continue
}
rm, err := MatcherFromConstraints(raid.GetConstraints())
if err != nil {
return nil, trace.Wrap(err)
}
if rm != nil {
matchers = append(matchers, rm)
}
}
return matchers, nil
}
// MatcherFromConstraints constructs a RoleMatcher encoding the requested
// ResourceConstraints for role resolution/validation time.
//
// This matcher is intended for use in request expansion, to decide whether a
// role qualifies for a resource where ResourceConstraints are specified.
//
// For enforcement of ResourceConstraints at authorization time, use
// WithConstraints to decorate principal-bearing matchers instead.
func MatcherFromConstraints(rc *types.ResourceConstraints) (RoleMatcher, error) {
if rc == nil {
return nil, nil
}
switch d := rc.Details.(type) {
case *types.ResourceConstraints_AwsConsole:
if err := d.Validate(); err != nil {
return nil, trace.Wrap(err)
}
matchers := make([]RoleMatcher, 0, len(d.AwsConsole.RoleArns))
for _, arn := range d.AwsConsole.RoleArns {
matchers = append(matchers, &AWSRoleARNMatcher{RoleARN: arn})
}
return RoleMatchers(matchers).AnyOf(), nil
case *types.ResourceConstraints_Ssh:
if err := d.Validate(); err != nil {
return nil, trace.Wrap(err)
}
matchers := make([]RoleMatcher, 0, len(d.Ssh.Logins))
for _, login := range d.Ssh.Logins {
matchers = append(matchers, NewLoginMatcher(login))
}
return RoleMatchers(matchers).AnyOf(), nil
default:
return nil, trace.BadParameter("unsupported constraint details type %T", d)
}
}
/*
* 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"
"encoding/json"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
type Restrictions interface {
GetNetworkRestrictions(context.Context) (types.NetworkRestrictions, error)
SetNetworkRestrictions(context.Context, types.NetworkRestrictions) error
DeleteNetworkRestrictions(context.Context) error
}
// ValidateNetworkRestrictions validates the network restrictions and sets defaults
func ValidateNetworkRestrictions(nr *types.NetworkRestrictionsV4) error {
if err := nr.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
return nil
}
// UnmarshalReverseTunnel unmarshals the ReverseTunnel resource from JSON.
func UnmarshalNetworkRestrictions(bytes []byte, opts ...MarshalOption) (types.NetworkRestrictions, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing network restrictions data")
}
var h types.ResourceHeader
err := json.Unmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V4:
var nr types.NetworkRestrictionsV4
if err := utils.FastUnmarshal(bytes, &nr); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := ValidateNetworkRestrictions(&nr); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
nr.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
nr.SetExpiry(cfg.Expires)
}
return &nr, nil
}
return nil, trace.BadParameter("network restrictions version %v is not supported", h.Version)
}
// MarshalNetworkRestrictions marshals the NetworkRestrictions resource to JSON.
func MarshalNetworkRestrictions(restrictions types.NetworkRestrictions, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if version := restrictions.GetVersion(); version != types.V4 {
return nil, trace.BadParameter("mismatched network restrictions version %v and type %T", version, restrictions)
}
switch restrictions := restrictions.(type) {
case *types.NetworkRestrictionsV4:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, restrictions))
default:
return nil, trace.BadParameter("unrecognized network restrictions version %T", restrictions)
}
}
/*
* 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 (
"bytes"
"cmp"
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net"
"os"
"path"
"reflect"
"regexp"
"slices"
"sort"
"strings"
"time"
"github.com/aws/aws-sdk-go-v2/aws/arn"
"github.com/google/uuid"
"github.com/gravitational/trace"
jsoniter "github.com/json-iterator/go"
"github.com/vulcand/predicate"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/defaults"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/wrappers"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/keys"
dtauthz "github.com/gravitational/teleport/lib/devicetrust/authz"
"github.com/gravitational/teleport/lib/services/label"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
awsutils "github.com/gravitational/teleport/lib/utils/aws"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/parse"
setutils "github.com/gravitational/teleport/lib/utils/set"
)
// DefaultImplicitRules provides access to the default set of implicit rules
// assigned to all roles.
var DefaultImplicitRules = []types.Rule{
types.NewRule(types.KindNode, RO()),
types.NewRule(types.KindProxy, RO()),
types.NewRule(types.KindAuthServer, RO()),
types.NewRule(types.KindReverseTunnel, RO()),
types.NewRule(types.KindCertAuthority, ReadNoSecrets()),
types.NewRule(types.KindClusterAuthPreference, RO()),
types.NewRule(types.KindClusterName, RO()),
types.NewRule(types.KindSSHSession, RO()),
types.NewRule(types.KindAppServer, RO()),
types.NewRule(types.KindRemoteCluster, RO()),
types.NewRule(types.KindKubeServer, RO()),
types.NewRule(types.KindDatabaseServer, RO()),
types.NewRule(types.KindDatabase, RO()),
types.NewRule(types.KindApp, RO()),
types.NewRule(types.KindWindowsDesktopService, RO()),
types.NewRule(types.KindWindowsDesktop, RO()),
types.NewRule(types.KindDynamicWindowsDesktop, RO()),
types.NewRule(types.KindLinuxDesktop, RO()),
types.NewRule(types.KindKubernetesCluster, RO()),
types.NewRule(types.KindUsageEvent, []string{types.VerbCreate}),
types.NewRule(types.KindVnetConfig, RO()),
types.NewRule(types.KindSPIFFEFederation, RO()),
types.NewRule(types.KindSAMLIdPServiceProvider, RO()),
types.NewRule(types.KindIdentityCenter, RO()),
types.NewRule(types.KindGitServer, RO()),
}
// DefaultCertAuthorityRules provides access the minimal set of resources
// needed for a certificate authority to function.
var DefaultCertAuthorityRules = []types.Rule{
types.NewRule(types.KindSession, RO()),
types.NewRule(types.KindNode, RO()),
types.NewRule(types.KindAuthServer, RO()),
types.NewRule(types.KindReverseTunnel, RO()),
types.NewRule(types.KindCertAuthority, ReadNoSecrets()),
}
// ErrTrustedDeviceRequired is returned by AccessChecker when access to a
// resource requires a trusted device.
// It's an alias to [dtauthz.ErrTrustedDeviceRequired].
var ErrTrustedDeviceRequired = dtauthz.ErrTrustedDeviceRequired
// ErrSessionMFARequired is returned by AccessChecker when access to a resource
// requires an MFA check.
var ErrSessionMFARequired = &trace.AccessDeniedError{
Message: "access to resource requires MFA",
}
// ErrSessionMFANotRequired indicates that per session mfa will not grant
// access to a resource.
var ErrSessionMFANotRequired = &trace.AccessDeniedError{
Message: "MFA is not required to access resource",
}
// RoleNameForUser returns role name associated with a user.
func RoleNameForUser(name string) string {
return "user:" + name
}
// RoleNameForCertAuthority returns role name associated with a certificate
// authority.
func RoleNameForCertAuthority(name string) string {
return "ca:" + name
}
// NewImplicitRole is the default implicit role that gets added to all
// RoleSets.
func NewImplicitRole() types.Role {
return newImplicitRole(types.CopyRulesSlice(DefaultImplicitRules))
}
// newScopedImplicitRole is the default implicit role as it applies to scoped identities. It confers
// the same privileges as [NewImplicitRole], except that secret-inclusive read is replaced with
// secret-exclusive read.
func newScopedImplicitRole() types.Role {
rules := types.CopyRulesSlice(DefaultImplicitRules)
for i, rule := range rules {
// CopyRulesSlice is a shallow copy, so build a replacement slice rather than assigning into it.
verbs := make([]string, len(rule.Verbs))
for j, verb := range rule.Verbs {
if verb == types.VerbRead {
verb = types.VerbReadNoSecrets
}
verbs[j] = verb
}
rules[i].Verbs = verbs
}
return newImplicitRole(rules)
}
func newImplicitRole(rules []types.Rule) types.Role {
return &types.RoleV6{
Kind: types.KindRole,
Version: types.V3,
Metadata: types.Metadata{
Name: constants.DefaultImplicitRole,
Namespace: defaults.Namespace,
},
Spec: types.RoleSpecV6{
Options: types.RoleOptions{
MaxSessionTTL: types.MaxDuration(),
RecordSession: &types.RecordSession{
Desktop: types.NewBoolOption(false),
},
},
Allow: types.RoleConditions{
Namespaces: []string{defaults.Namespace},
Rules: rules,
},
},
}
}
// RoleForUser creates an admin role for a services.User.
//
// Used in tests only.
func RoleForUser(u types.User) types.Role {
return RoleWithVersionForUser(u, types.DefaultRoleVersion)
}
// RoleWithVersionForUser creates an admin role for a services.User.
//
// Used in tests only.
func RoleWithVersionForUser(u types.User, v string) types.Role {
role, _ := types.NewRoleWithVersion(RoleNameForUser(u.GetName()), v, types.RoleSpecV6{
Options: types.RoleOptions{
CertificateFormat: constants.CertificateFormatStandard,
MaxSessionTTL: types.NewDuration(defaults.MaxCertDuration),
PortForwarding: types.NewBoolOption(true),
ForwardAgent: types.NewBool(true),
BPF: defaults.EnhancedEvents(),
},
Allow: types.RoleConditions{
Namespaces: []string{defaults.Namespace},
NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
AppLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
GroupLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
KubernetesLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseServiceLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
MCP: &types.MCPPermissions{
Tools: []string{types.Wildcard},
},
Rules: []types.Rule{
types.NewRule(types.KindRole, RW()),
types.NewRule(types.KindAuthConnector, RW()),
types.NewRule(types.KindSession, RO()),
types.NewRule(types.KindTrustedCluster, RW()),
types.NewRule(types.KindEvent, RO()),
types.NewRule(types.KindClusterAuthPreference, RW()),
types.NewRule(types.KindClusterNetworkingConfig, RW()),
types.NewRule(types.KindSessionRecordingConfig, RW()),
types.NewRule(types.KindUIConfig, RW()),
types.NewRule(types.KindApp, RW()),
types.NewRule(types.KindDatabase, RW()),
types.NewRule(types.KindLock, RW()),
types.NewRule(types.KindToken, RW()),
types.NewRule(types.KindConnectionDiagnostic, RW()),
types.NewRule(types.KindKubernetesCluster, RW()),
types.NewRule(types.KindSessionTracker, RO()),
types.NewRule(types.KindUserGroup, RW()),
types.NewRule(types.KindSAMLIdPServiceProvider, RW()),
},
JoinSessions: []*types.SessionJoinPolicy{
{
Name: "foo",
Roles: []string{"*"},
Kinds: []string{string(types.SSHSessionKind)},
Modes: []string{string(types.SessionPeerMode)},
},
},
},
})
return role
}
// RoleForCertAuthority creates role using types.CertAuthority.
func RoleForCertAuthority(ca types.CertAuthority) types.Role {
role, _ := types.NewRole(RoleNameForCertAuthority(ca.GetClusterName()), types.RoleSpecV6{
Options: types.RoleOptions{
MaxSessionTTL: types.NewDuration(defaults.MaxCertDuration),
},
Allow: types.RoleConditions{
Namespaces: []string{defaults.Namespace},
NodeLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
AppLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
KubernetesLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
DatabaseLabels: types.Labels{types.Wildcard: []string{types.Wildcard}},
Rules: types.CopyRulesSlice(DefaultCertAuthorityRules),
},
})
return role
}
// ValidateRoleName checks that the role name is allowed to be created.
func ValidateRoleName(role types.Role) error {
// System role names are not allowed.
if types.SystemRole(role.GetMetadata().Name).IsValid() {
return trace.BadParameter("reserved role: %+q", role.GetMetadata().Name)
}
return nil
}
// ValidateRole checks and sets defaults for role fields and validates
// expression syntax.
//
// This function should be called on the write path (role create/update)
// and NOT on read paths to avoid bricking clusters with existing roles
// that may not parse with newer parsers. Read paths should call plain
// CheckAndSetDefaults to reject truly unusable roles.
func ValidateRole(r types.Role) error {
if err := CheckAndSetDefaults(r); err != nil {
return trace.Wrap(err)
}
var errs []error
if err := validateRoleExpressions(r); err != nil {
errs = append(errs, err)
}
if err := validateRoleWildcards(r); err != nil {
errs = append(errs, err)
}
if err := validateSessionPolicies(r); err != nil {
errs = append(errs, err)
}
if err := validateAppResources(r); err != nil {
errs = append(errs, err)
}
return trace.NewAggregate(errs...)
}
// validateAppResources rejects an app_resources rule set that this version
// cannot enforce, for example a rule with an unknown field. It also rejects
// any app_resources_expressions. It runs on create and update only, not on
// read.
func validateAppResources(r types.Role) error {
if len(r.GetAppResources(types.Deny)) > 0 {
return trace.BadParameter("app_resources is not allowed under deny")
}
if len(r.GetAppResourcesExpressions(types.Deny)) > 0 {
return trace.BadParameter("app_resources_expressions is not allowed under deny")
}
if len(r.GetAppResourcesExpressions(types.Allow)) > 0 {
return trace.BadParameter("app_resources_expressions is not supported in this version, only app_resources with allow_all is honored")
}
allow := r.GetAppResources(types.Allow)
for i, rule := range allow {
// The backend JSON marshal drops unknown fields. Storing such a
// rule would silently widen it to unrestricted access.
if !rule.IsAllowAllOnly() {
return trace.BadParameter("app_resources[%d]: this version implements allow_all only, so a rule must set allow_all and nothing else", i)
}
}
// Every rule sets allow_all at this point, so more than one rule can
// only mean allow_all next to another rule.
if len(allow) > 1 {
return trace.BadParameter("app_resources: a rule setting allow_all must be the only rule")
}
return nil
}
// validateRoleExpressions validates all expression and predicate syntax in a role.
func validateRoleExpressions(r types.Role) error {
var errs []error
for _, condition := range []struct {
name string
condition types.RoleConditionType
}{
{"allow", types.Allow},
{"deny", types.Deny},
} {
// Rules
for i, rule := range r.GetRules(condition.condition) {
if err := validateRule(rule); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.rules[%d]: %v", condition.name, i, err))
}
}
// Trait templates in slice fields
for _, values := range []struct {
name string
values []string
}{
{"logins", r.GetLogins(condition.condition)},
{"windows_desktop_logins", r.GetWindowsLogins(condition.condition)},
{"linux_desktop_logins", r.GetLinuxDesktopLogins(condition.condition)},
{"aws_role_arns", r.GetAWSRoleARNs(condition.condition)},
{"azure_identities", r.GetAzureIdentities(condition.condition)},
{"gcp_service_accounts", r.GetGCPServiceAccounts(condition.condition)},
{"kubernetes_groups", r.GetKubeGroups(condition.condition)},
{"kubernetes_users", r.GetKubeUsers(condition.condition)},
{"db_names", r.GetDatabaseNames(condition.condition)},
{"db_users", r.GetDatabaseUsers(condition.condition)},
{"db_roles", r.GetDatabaseRoles(condition.condition)},
{"host_groups", r.GetHostGroups(condition.condition)},
{"host_sudoers", r.GetHostSudoers(condition.condition)},
{"desktop_groups", r.GetDesktopGroups(condition.condition)},
{"impersonate.users", r.GetImpersonateConditions(condition.condition).Users},
{"impersonate.roles", r.GetImpersonateConditions(condition.condition).Roles},
} {
for _, value := range values.values {
if _, err := parse.NewTraitsTemplateExpression(value); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.%s expression: %v", condition.name, values.name, err))
}
}
}
// Impersonate where clause
if where := r.GetImpersonateConditions(condition.condition).Where; where != "" {
// Stub the context, the predicate parser resolves identifiers at parse
// time, so a nil context rejects valid expressions.
parser, err := newImpersonateWhereParser(&impersonateContext{
user: emptyUser,
impersonateUser: emptyUser,
impersonateRole: &types.RoleV6{},
})
if err != nil {
errs = append(errs, trace.BadParameter("%s.impersonate.where: failed to create parser: %v", condition.name, err))
} else if _, err = parser.Parse(where); err != nil {
errs = append(errs, trace.BadParameter("%s.impersonate.where: invalid expression %q: %v", condition.name, where, err))
}
}
// Trait templates in kubernetes_resources
for i, ks := range r.GetKubeResources(condition.condition) {
if _, err := parse.NewTraitsTemplateExpression(ks.Namespace); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.kubernetes_resources[%d].namespace expression: %v", condition.name, i, err))
}
if _, err := parse.NewTraitsTemplateExpression(ks.Name); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.kubernetes_resources[%d].name expression: %v", condition.name, i, err))
}
for _, verb := range ks.Verbs {
if _, err := parse.NewTraitsTemplateExpression(verb); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.kubernetes_resources[%d].verbs expression: %v", condition.name, i, err))
}
}
}
// Label value trait templates and label expressions
for _, labels := range []struct {
name string
kind string
}{
{"cluster_labels", types.KindRemoteCluster},
{"node_labels", types.KindNode},
{"kubernetes_labels", types.KindKubernetesCluster},
{"app_labels", types.KindApp},
{"saml_idp_service_provider", types.KindSAMLIdPServiceProvider},
{"db_labels", types.KindDatabase},
{"db_service_labels", types.KindDatabaseService},
{"windows_desktop_labels", types.KindWindowsDesktop},
{"windows_desktop_labels", types.KindDynamicWindowsDesktop},
{"group_labels", types.KindUserGroup},
{"workload_identity_labels", types.KindWorkloadIdentity},
{"beam_labels", types.KindBeam},
} {
labelMatchers, err := r.GetLabelMatchers(condition.condition, labels.kind)
if err != nil {
return trace.Wrap(err)
}
for _, labelValues := range labelMatchers.Labels {
for _, label := range labelValues {
if _, err := parse.NewTraitsTemplateExpression(label); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.%s template expression: %v", condition.name, labels.name, err))
}
}
}
if len(labelMatchers.Expression) > 0 {
if _, err := label.ParseExpression(labelMatchers.Expression); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.%s_expression: %v", condition.name, labels.name, err))
}
}
}
// Trait templates in github_permissions.organizations
for i, perm := range r.GetGitHubPermissions(condition.condition) {
for _, org := range perm.Organizations {
if _, err := parse.NewTraitsTemplateExpression(org); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.github_permissions[%d].organizations expression: %v", condition.name, i, err))
}
}
}
// Trait templates in mcp.tools
if mcp := r.GetMCPPermissions(condition.condition); mcp != nil {
for i, tool := range mcp.Tools {
if _, err := parse.NewTraitsTemplateExpression(tool); err != nil {
errs = append(errs, trace.BadParameter("parsing %s.mcp.tools[%d] %q: %v", condition.name, i, tool, err))
}
}
}
}
// Trait templates in options.cert_extensions.value
for i, ext := range r.GetOptions().CertExtensions {
if ext == nil {
continue
}
if _, err := parse.NewTraitsTemplateExpression(ext.Value); err != nil {
errs = append(errs, trace.BadParameter("parsing options.cert_extensions[%d].value expression: %v", i, err))
}
}
// Session require policy expressions
for i, policy := range r.GetSessionRequirePolicies() {
if policy == nil || policy.Filter == "" {
continue
}
parser, err := NewWhereParser(sessionFilterValidationContext{})
if err != nil {
errs = append(errs, trace.BadParameter("require_session_join[%d]: failed to create where parser: %v", i, err))
continue
}
if _, err = parser.Parse(policy.Filter); err != nil {
errs = append(errs, trace.BadParameter("require_session_join[%d]: invalid filter %q: %v", i, policy.Filter, err))
}
}
// Access predicates
if err := ValidateAccessPredicates(r); err != nil {
errs = append(errs, err)
}
return trace.NewAggregate(errs...)
}
// sessionFilterValidationContext validates require_session_join filters using
// the same identifiers as moderation.SessionAccessContext at runtime including
// the legacy user.roles alias for user.spec.roles.
//
// Keep in sync with moderation.SessionAccessContext.GetIdentifier in
// lib/auth/moderation/session_access.go.
type sessionFilterValidationContext struct{}
func (sessionFilterValidationContext) GetIdentifier(fields []string) (any, error) {
if fields[0] == "user" && (len(fields) == 2 || len(fields) == 3) {
idx := 1
if len(fields) == 3 && fields[1] == "spec" {
idx = 2
}
switch fields[idx] {
case "name":
return "", nil
case "roles":
return []string{}, nil
}
}
return nil, trace.NotFound("%v is not defined", strings.Join(fields, "."))
}
func (sessionFilterValidationContext) GetResource() (types.Resource, error) {
return nil, trace.NotFound("resource is not used in session filter validation")
}
func (sessionFilterValidationContext) GetAccessChecker() (AccessChecker, error) {
return nil, trace.NotFound("access checker is not used in session filter validation")
}
// validateRoleWildcards rejects wildcards in fields that don't support them.
func validateRoleWildcards(r types.Role) error {
var errs []error
for _, side := range []struct {
name string
rct types.RoleConditionType
}{
{"allow", types.Allow},
{"deny", types.Deny},
} {
for _, field := range []struct {
name string
values []string
}{
{"request.search_as_roles", r.GetSearchAsRoles(side.rct)},
{"review_requests.preview_as_roles", r.GetPreviewAsRoles(side.rct)},
} {
if slices.Contains(field.values, types.Wildcard) {
errs = append(errs, trace.BadParameter("wildcard is not allowed in %s.%s", side.name, field.name))
}
}
}
return trace.NewAggregate(errs...)
}
// validateSessionPolicies validates fields (kinds, modes,
// on_leave) on require_session_join join_sessions policies.
func validateSessionPolicies(r types.Role) error {
var errs []error
// require_session_join
for i, p := range r.GetSessionRequirePolicies() {
if p == nil {
continue
}
if p.Count < 0 {
errs = append(errs, trace.BadParameter("require_session_join[%d]: count cannot be negative, got %d", i, p.Count))
}
if err := validateSessionKinds(p.Kinds); err != nil {
errs = append(errs, trace.BadParameter("require_session_join[%d]: %v", i, err))
}
if err := validateSessionParticipantModes(p.Modes); err != nil {
errs = append(errs, trace.BadParameter("require_session_join[%d]: %v", i, err))
}
switch types.OnSessionLeaveAction(p.OnLeave) {
case "", types.OnSessionLeaveTerminate, types.OnSessionLeavePause:
default:
errs = append(errs, trace.BadParameter("require_session_join[%d]: invalid on_leave action %q, expected one of %q, %q",
i, p.OnLeave, types.OnSessionLeaveTerminate, types.OnSessionLeavePause))
}
}
// join_sessions
for i, p := range r.GetSessionJoinPolicies() {
if p == nil {
continue
}
if err := validateSessionKinds(p.Kinds); err != nil {
errs = append(errs, trace.BadParameter("join_sessions[%d]: %v", i, err))
}
if err := validateSessionParticipantModes(p.Modes); err != nil {
errs = append(errs, trace.BadParameter("join_sessions[%d]: %v", i, err))
}
}
return trace.NewAggregate(errs...)
}
func validateSessionKinds(kinds []string) error {
for _, kind := range kinds {
// "*" is accepted by SessionAccessEvaluator.matchesKind at runtime
if kind == types.Wildcard {
continue
}
switch types.SessionKind(kind) {
case types.SSHSessionKind, types.KubernetesSessionKind, types.DatabaseSessionKind,
types.AppSessionKind, types.WindowsDesktopSessionKind, types.GitSessionKind:
default:
return trace.BadParameter("invalid session kind %q", kind)
}
}
return nil
}
func validateSessionParticipantModes(modes []string) error {
for _, mode := range modes {
switch types.SessionParticipantMode(mode) {
case types.SessionObserverMode, types.SessionModeratorMode, types.SessionPeerMode:
default:
return trace.BadParameter("invalid participant mode %q", mode)
}
}
return nil
}
// validateRule parses the where and action fields to validate the rule.
func validateRule(r types.Rule) error {
if len(r.Where) != 0 {
parser, err := NewWhereParser(&Context{},
ConditionalOption(
slices.Contains(r.Resources, types.KindSession),
WithCanViewFunction(),
),
)
if err != nil {
return trace.Wrap(err)
}
_, err = parser.Parse(r.Where)
if err != nil {
return trace.BadParameter("could not parse 'where' rule: %q, error: %v", r.Where, err)
}
}
if len(r.Actions) != 0 {
parser, err := NewActionsParser(&Context{})
if err != nil {
return trace.Wrap(err)
}
for i, action := range r.Actions {
_, err = parser.Parse(action)
if err != nil {
return trace.BadParameter("could not parse action %v %q, error: %v", i, action, err)
}
}
}
return nil
}
func filterInvalidUnixLogins(candidates []string) []string {
var output []string
for _, candidate := range candidates {
if utils.IsValidUnixUser(candidate) {
// A valid variable was found in the traits, append it to the list of logins.
output = append(output, candidate)
continue
}
// Log any invalid logins which were added by a user but ignore any
// Teleport internal logins which are known to be invalid.
if candidate != teleport.SSHSessionJoinPrincipal && !strings.HasPrefix(candidate, "no-login-") {
slog.DebugContext(context.Background(), "Skipping invalid Unix login.", "login", candidate)
}
}
return output
}
func filterInvalidWindowsLogins(candidates []string) []string {
var output []string
// https://docs.microsoft.com/en-us/previous-versions/windows/it-pro/windows-2000-server/bb726984(v=technet.10)
const invalidChars = `"/\[]:;|=,+*?<>`
for _, candidate := range candidates {
if strings.ContainsAny(candidate, invalidChars) {
slog.DebugContext(context.Background(), "Skipping invalid Windows login.", "login", candidate)
continue
}
output = append(output, candidate)
}
return output
}
func warnInvalidAzureIdentities(candidates []string) {
for _, candidate := range candidates {
if !MatchValidAzureIdentity(candidate) {
slog.WarnContext(context.Background(), "Invalid format of Azure identity", "identity", candidate)
}
}
}
// ParseResourceID from Azure SDK is too lenient; we use a strict regexp instead.
var azureIdentityPattern = regexp.MustCompile(`(?i)^/subscriptions/([a-fA-F0-9-]+)/resourceGroups/([0-9a-zA-Z-_]+)/providers/Microsoft\.ManagedIdentity/userAssignedIdentities/([0-9a-zA-Z-_]+)$`)
func MatchValidAzureIdentity(identity string) bool {
if identity == types.Wildcard {
return true
}
return azureIdentityPattern.MatchString(identity)
}
// RoleTemplateContext is the runtime context used to evaluate role template
// expressions.
type RoleTemplateContext struct {
Username string
Traits map[string][]string
}
// ApplyTraits applies the passed in traits to any variables within the role
// and returns itself.
func ApplyTraits(r types.Role, traits map[string][]string) (types.Role, error) {
return ApplyTraitsWithContext(r, RoleTemplateContext{Traits: traits})
}
// ApplyTraitsWithContext applies the passed in role template context to any
// variables within the role and returns itself.
//
// Keep in sync with validateRoleExpressions.
func ApplyTraitsWithContext(r types.Role, ctx RoleTemplateContext) (types.Role, error) {
for _, condition := range []types.RoleConditionType{types.Allow, types.Deny} {
inLogins := r.GetLogins(condition)
outLogins := applyValueTraitsSlice(inLogins, ctx, "login")
outLogins = filterInvalidUnixLogins(outLogins)
r.SetLogins(condition, apiutils.Deduplicate(outLogins))
inWindowsLogins := r.GetWindowsLogins(condition)
outWindowsLogins := applyValueTraitsSlice(inWindowsLogins, ctx, "windows_login")
outWindowsLogins = filterInvalidWindowsLogins(outWindowsLogins)
r.SetWindowsLogins(condition, apiutils.Deduplicate(outWindowsLogins))
inLinuxDesktopLogins := r.GetLinuxDesktopLogins(condition)
outLinuxDesktopLogins := applyValueTraitsSlice(inLinuxDesktopLogins, ctx, "linux_desktop_login")
outLinuxDesktopLogins = filterInvalidUnixLogins(outLinuxDesktopLogins)
r.SetLinuxDesktopLogins(condition, apiutils.Deduplicate(outLinuxDesktopLogins))
inRoleARNs := r.GetAWSRoleARNs(condition)
outRoleARNs := applyValueTraitsSlice(inRoleARNs, ctx, "AWS role ARN")
r.SetAWSRoleARNs(condition, apiutils.Deduplicate(outRoleARNs))
inAzureIdentities := r.GetAzureIdentities(condition)
outAzureIdentities := applyValueTraitsSlice(inAzureIdentities, ctx, "Azure identity")
warnInvalidAzureIdentities(outAzureIdentities)
r.SetAzureIdentities(condition, apiutils.Deduplicate(outAzureIdentities))
inGCPAccounts := r.GetGCPServiceAccounts(condition)
outGCPAccounts := applyValueTraitsSlice(inGCPAccounts, ctx, "GCP service account")
r.SetGCPServiceAccounts(condition, apiutils.Deduplicate(outGCPAccounts))
// apply templates to kubernetes groups
inKubeGroups := r.GetKubeGroups(condition)
outKubeGroups := applyValueTraitsSlice(inKubeGroups, ctx, "kube group")
r.SetKubeGroups(condition, apiutils.Deduplicate(outKubeGroups))
// apply templates to kubernetes users
inKubeUsers := r.GetKubeUsers(condition)
outKubeUsers := applyValueTraitsSlice(inKubeUsers, ctx, "kube user")
r.SetKubeUsers(condition, apiutils.Deduplicate(outKubeUsers))
// apply templates to database names
inDbNames := r.GetDatabaseNames(condition)
outDbNames := applyValueTraitsSlice(inDbNames, ctx, "database name")
r.SetDatabaseNames(condition, apiutils.Deduplicate(outDbNames))
// apply templates to database users
inDbUsers := r.GetDatabaseUsers(condition)
outDbUsers := applyValueTraitsSlice(inDbUsers, ctx, "database user")
r.SetDatabaseUsers(condition, apiutils.Deduplicate(outDbUsers))
// apply templates to database roles
inDbRoles := r.GetDatabaseRoles(condition)
outDbRoles := applyValueTraitsSlice(inDbRoles, ctx, "database role")
r.SetDatabaseRoles(condition, apiutils.Deduplicate(outDbRoles))
githubPermissions := r.GetGitHubPermissions(condition)
for i, perm := range githubPermissions {
githubPermissions[i].Organizations = applyValueTraitsSlice(perm.Organizations, ctx, "github organizations")
}
r.SetGitHubPermissions(condition, githubPermissions)
var out []types.KubernetesResource
// we access the resources in the role using the role conditions
// to avoid receiving the compatibility resources added in GetKubernetesResources
// for roles <v7
for _, rec := range r.GetRoleConditions(condition).KubernetesResources {
namespaces := applyValueTraitsSlice([]string{rec.Namespace}, ctx, "kubernetes resource namespace")
if rec.Namespace == "" {
namespaces = []string{""}
}
names := applyValueTraitsSlice([]string{rec.Name}, ctx, "kubernetes resource name")
if rec.Name == "" {
names = []string{""}
}
verbs := applyValueTraitsSlice(rec.Verbs, ctx, "kubernetes resource verb")
// A trait can reintroduce a wildcard alongside other verbs after
// validation has run, so collapse to just the wildcard.
if slices.Contains(verbs, types.Wildcard) {
verbs = []string{types.Wildcard}
}
for _, namespace := range namespaces {
for _, name := range names {
out = append(out, types.KubernetesResource{
Kind: rec.Kind,
Namespace: namespace,
Name: name,
Verbs: verbs,
APIGroup: rec.APIGroup,
})
}
}
}
r.SetKubeResources(condition, out)
for _, kind := range []string{
types.KindRemoteCluster,
types.KindNode,
types.KindKubernetesCluster,
types.KindApp,
types.KindDatabase,
types.KindDatabaseService,
types.KindWindowsDesktop,
types.KindLinuxDesktop,
types.KindUserGroup,
types.KindSAMLIdPServiceProvider,
types.KindWorkloadIdentity,
types.KindBeam,
} {
labelMatchers, err := r.GetLabelMatchers(condition, kind)
if err != nil {
return nil, trace.Wrap(err)
}
// Only labelMatchers.Labels is templated, if empty we can skip
// these label matchers. labelMatchers.Expression can reference user
// traits later during the access check through the expression
// environment, they are not templated in here.
if len(labelMatchers.Labels) == 0 {
continue
}
labelMatchers.Labels = applyLabelsTraits(labelMatchers.Labels, ctx)
if err := r.SetLabelMatchers(condition, kind, labelMatchers); err != nil {
return nil, trace.Wrap(err)
}
}
r.SetHostGroups(condition,
applyValueTraitsSlice(r.GetHostGroups(condition), ctx, "host_groups"))
r.SetHostSudoers(condition,
applyValueTraitsSlice(r.GetHostSudoers(condition), ctx, "host_sudoers"))
r.SetDesktopGroups(condition,
applyValueTraitsSlice(r.GetDesktopGroups(condition), ctx, "desktop_groups"))
options := r.GetOptions()
for i, ext := range options.CertExtensions {
vals, err := ApplyValueTraitsWithContext(ext.Value, ctx)
if err != nil && !trace.IsNotFound(err) {
slog.WarnContext(context.Background(), "Failed to apply trait to cert_extensions.value", "error", err)
continue
}
if len(vals) != 0 {
options.CertExtensions[i].Value = vals[0]
}
}
// apply templates to impersonation conditions
inCond := r.GetImpersonateConditions(condition)
var outCond types.ImpersonateConditions
outCond.Users = applyValueTraitsSlice(inCond.Users, ctx, "impersonate user")
outCond.Roles = applyValueTraitsSlice(inCond.Roles, ctx, "impersonate role")
outCond.Users = apiutils.Deduplicate(outCond.Users)
outCond.Roles = apiutils.Deduplicate(outCond.Roles)
outCond.Where = inCond.Where
r.SetImpersonateConditions(condition, outCond)
if mcp := r.GetMCPPermissions(condition); mcp != nil {
mcp.Tools = applyValueTraitsSlice(mcp.Tools, ctx, "mcp.tools")
r.SetMCPPermissions(condition, mcp)
}
}
return r, nil
}
// applyValueTraitsSlice iterates over a slice of input strings, calling
// ApplyValueTraitsWithContext on each.
func applyValueTraitsSlice(inputs []string, ctx RoleTemplateContext, fieldName string) []string {
var output []string
for _, value := range inputs {
outputs, err := ApplyValueTraitsWithContext(value, ctx)
if err != nil {
if !trace.IsNotFound(err) {
slog.DebugContext(context.Background(), "Skipping trait value.", "field", fieldName, "value", value, "error", err)
}
continue
}
output = append(output, outputs...)
}
return output
}
// applyLabelsTraits interpolates variables based on the templates
// and traits from identity provider. For example:
//
// cluster_labels:
//
// env: ['{{external.groups}}']
//
// and groups: ['admins', 'devs']
//
// will be interpolated to:
//
// cluster_labels:
//
// env: ['admins', 'devs']
func applyLabelsTraits(inLabels types.Labels, ctx RoleTemplateContext) types.Labels {
outLabels := make(types.Labels, len(inLabels))
// every key will be mapped to the first value
for key, vals := range inLabels {
keyVars, err := ApplyValueTraitsWithContext(key, ctx)
if err != nil {
// empty key will not match anything
slog.DebugContext(context.Background(), "Setting empty node label pair", "key", key, "values", vals, "error", err)
keyVars = []string{""}
}
var values []string
for _, val := range vals {
valVars, err := ApplyValueTraitsWithContext(val, ctx)
if err != nil {
slog.DebugContext(context.Background(), "Setting empty node label value", "key", key, "value", val, "error", err)
// empty value will not match anything
valVars = []string{""}
}
values = append(values, valVars...)
}
outLabels[keyVars[0]] = apiutils.Deduplicate(values)
}
return outLabels
}
// ApplyValueTraits applies the passed in traits to the variable,
// returns BadParameter in case if referenced variable is unsupported,
// returns NotFound in case if referenced trait is missing,
// mapped list of values otherwise, the function guarantees to return
// at least one value in case if return value is nil
func ApplyValueTraits(val string, traits map[string][]string) ([]string, error) {
return ApplyValueTraitsWithContext(val, RoleTemplateContext{Traits: traits})
}
// ApplyValueTraitsWithContext applies the passed in role template context to
// the variable.
func ApplyValueTraitsWithContext(val string, ctx RoleTemplateContext) ([]string, error) {
// Extract the variable from the role variable.
expr, err := parse.NewTraitsTemplateExpression(val)
if err != nil {
return nil, trace.Wrap(err)
}
varValidation := func(namespace string, name string) error {
// verify that internal traits match the supported variables
if namespace == teleport.TraitInternalPrefix {
switch name {
case constants.TraitLogins, constants.TraitWindowsLogins, constants.TraitLinuxDesktopLogins,
constants.TraitKubeGroups, constants.TraitKubeUsers,
constants.TraitDBNames, constants.TraitDBUsers, constants.TraitDBRoles,
constants.TraitAWSRoleARNs, constants.TraitAzureIdentities,
constants.TraitGCPServiceAccounts, constants.TraitJWT,
constants.TraitGitHubOrgs, constants.TraitMCPTools,
constants.TraitDefaultRelayAddr, constants.TraitIDToken:
default:
return trace.BadParameter("unsupported variable %q", name)
}
}
// The "external" trait namespace is explicitly allowed to reference
// "internal" traits listed above. This is for multiple reasons:
// - back compat, it's always been this way
// - IdPs are allowed to set those trait names so it wouldn't make
// sense to block them when referenced via "external"
// - The user resource spec.traits can include the "internal" trait
// names listed above, as well as any other trait name - but other
// trait names must be referenced in the "external" namespace. It
// wouldn't make a lot of sense to change the namespace
// based only on the trait name, especially given that we tend to
// expand the list of "internal" traits, and that would be a breaking
// change if someone already referred to one of the "new" internal
// traits in the "external" namespace.
return nil
}
interpolated, err := expr.InterpolateWithUser(varValidation, ctx.Username, ctx.Traits)
if err != nil {
return nil, trace.Wrap(err)
}
if len(interpolated) == 0 {
return nil, trace.NotFound("variable interpolation result is empty")
}
return interpolated, nil
}
// ruleScore is a sorting score of the rule, the larger the score, the more
// specific the rule is
func ruleScore(r *types.Rule) int {
score := 0
// wildcard rules are less specific
if slices.Contains(r.Resources, types.Wildcard) {
score -= 4
} else if len(r.Resources) == 1 {
// rules that match specific resource are more specific than
// fields that match several resources
score += 2
}
// rules that have wildcard verbs are less specific
if slices.Contains(r.Verbs, types.Wildcard) {
score -= 2
}
// rules that supply 'where' or 'actions' are more specific
// having 'where' or 'actions' is more important than
// whether the rules are wildcard or not, so here we have +8 vs
// -4 and -2 score penalty for wildcards in resources and verbs
if len(r.Where) > 0 {
score += 8
}
// rules featuring actions are more specific
if len(r.Actions) > 0 {
score += 8
}
return score
}
// CompareRuleScore returns true if the first rule is more specific than the other.
//
// * nRule matching wildcard resource is less specific
// than same rule matching specific resource.
// * Rule that has wildcard verbs is less specific
// than the same rules matching specific verb.
// * Rule that has where section is more specific
// than the same rule without where section.
// * Rule that has actions list is more specific than
// rule without actions list.
func CompareRuleScore(r *types.Rule, o *types.Rule) bool {
return ruleScore(r) > ruleScore(o)
}
// RuleSet maps resource to a set of rules defined for it
type RuleSet map[string][]types.Rule
// MakeRuleSet creates a new rule set from a list
func MakeRuleSet(rules []types.Rule) RuleSet {
set := make(RuleSet)
for _, rule := range rules {
for _, resource := range rule.Resources {
set[resource] = append(set[resource], rule)
}
}
for resource := range set {
rules := set[resource]
// sort rules by most specific rule, the rule that has actions
// is more specific than the one that has no actions
sort.Slice(rules, func(i, j int) bool {
return CompareRuleScore(&rules[i], &rules[j])
})
set[resource] = rules
}
return set
}
// Match tests if the resource name and verb are in a given list of rules.
// More specific rules will be matched first. See Rule.IsMoreSpecificThan
// for exact specs on whether the rule is more or less specific.
//
// Specifying order solves the problem on having multiple rules, e.g. one wildcard
// rule can override more specific rules with 'where' sections that can have
// 'actions' lists with side effects that will not be triggered otherwise.
func (set RuleSet) Match(whereParser predicate.Parser, actionsParser predicate.Parser, resource string, verb string) (bool, error) {
// empty set matches nothing
if len(set) == 0 {
return false, nil
}
// check for matching resource by name
// the most specific rule should win
rules := set[resource]
for _, rule := range rules {
match, err := matchesWhere(&rule, whereParser)
if err != nil {
return false, trace.Wrap(err)
}
if match && (rule.HasVerb(types.Wildcard) || rule.HasVerb(verb)) {
if err := processActions(&rule, actionsParser); err != nil {
return true, trace.Wrap(err)
}
return true, nil
}
}
// check for wildcard resource matcher
for _, rule := range set[types.Wildcard] {
match, err := matchesWhere(&rule, whereParser)
if err != nil {
return false, trace.Wrap(err)
}
if match && (rule.HasVerb(types.Wildcard) || rule.HasVerb(verb)) {
if err := processActions(&rule, actionsParser); err != nil {
return true, trace.Wrap(err)
}
return true, nil
}
}
return false, nil
}
// matchesWhere returns true if Where rule matches.
// Empty Where block always matches.
func matchesWhere(r *types.Rule, parser predicate.Parser) (bool, error) {
if r.Where == "" {
return true, nil
}
ifn, err := parser.Parse(r.Where)
if err != nil {
return false, trace.Wrap(err)
}
fn, ok := ifn.(predicate.BoolPredicate)
if !ok {
return false, trace.BadParameter("invalid predicate type for where expression: %v", r.Where)
}
return fn(), nil
}
// processActions processes actions specified for this rule
func processActions(r *types.Rule, parser predicate.Parser) error {
for _, action := range r.Actions {
ifn, err := parser.Parse(action)
if err != nil {
return trace.Wrap(err)
}
fn, ok := ifn.(predicate.BoolPredicate)
if !ok {
return trace.BadParameter("invalid predicate type for action expression: %v", action)
}
fn()
}
return nil
}
// Slice returns slice from a set
func (set RuleSet) Slice() []types.Rule {
var out []types.Rule
for _, rules := range set {
out = append(out, rules...)
}
return out
}
// RoleFromSpec returns new Role created from spec
func RoleFromSpec(name string, spec types.RoleSpecV6) (types.Role, error) {
role, err := types.NewRole(name, spec)
return role, trace.Wrap(err)
}
// RoleSetFromSpec returns a new RoleSet from spec
func RoleSetFromSpec(name string, spec types.RoleSpecV6) (RoleSet, error) {
role, err := RoleFromSpec(name, spec)
if err != nil {
return nil, trace.Wrap(err)
}
return NewRoleSet(role), nil
}
// WO is a shortcut that returns create and update verbs, granting the ability
// to emit/write resources but not list, read, or delete them.
func WO() []string {
return []string{types.VerbCreate, types.VerbUpdate}
}
// RW is a shortcut that returns all CRUD verbs.
func RW() []string {
return []string{types.VerbList, types.VerbCreate, types.VerbRead, types.VerbUpdate, types.VerbDelete}
}
// RO is a shortcut that returns read only verbs that provide access to secrets.
func RO() []string {
return []string{types.VerbList, types.VerbRead}
}
// ReadNoSecrets is a shortcut that returns read only verbs that do not
// provide access to secrets.
func ReadNoSecrets() []string {
return []string{types.VerbList, types.VerbReadNoSecrets}
}
// RoleGetter is an interface that defines GetRole method
type RoleGetter interface {
// GetRole returns role by name
GetRole(ctx context.Context, name string) (types.Role, error)
}
// ExtractFromIdentity will extract roles and traits from the *x509.Certificate
// which Teleport passes along as a *tlsca.Identity. If roles and traits do not
// exist in the certificates, they are extracted from the backend.
func ExtractFromIdentity(ctx context.Context, access UserGetter, identity tlsca.Identity) ([]string, wrappers.Traits, error) {
// Legacy certs are not encoded with roles or traits,
// so we fallback to the traits and roles in the backend.
// empty traits are a valid use case in standard certs,
// so we only check for whether roles are empty.
if len(identity.Groups) == 0 {
u, err := access.GetUser(ctx, identity.Username, false)
if err != nil {
return nil, nil, trace.Wrap(err)
}
const msg = "Failed to find roles in x509 identity. Fetching " +
"from backend. If the identity provider allows username changes, this can " +
"potentially allow an attacker to change the role of the existing user."
slog.WarnContext(ctx, msg, "username", identity.Username)
return u.GetRoles(), u.GetTraits(), nil
}
return identity.Groups, identity.Traits, nil
}
// FetchRoleList fetches roles by their names, applies the traits to role
// variables, and returns the list
func FetchRoleList(roleNames []string, access RoleGetter, traits map[string][]string) (RoleSet, error) {
return FetchRoleListWithContext(roleNames, access, RoleTemplateContext{Traits: traits})
}
// FetchRoleListWithContext fetches roles by their names, applies the role
// template context to role variables, and returns the list.
func FetchRoleListWithContext(roleNames []string, access RoleGetter, ctx RoleTemplateContext) (RoleSet, error) {
var roles []types.Role
for _, roleName := range roleNames {
role, err := access.GetRole(context.TODO(), roleName)
if err != nil {
return nil, trace.Wrap(err)
}
role, err = ApplyTraitsWithContext(role, ctx)
if err != nil {
return nil, trace.Wrap(err)
}
roles = append(roles, role)
}
return roles, nil
}
// FetchRoles fetches roles by their names, applies the traits to role
// variables, and returns the RoleSet. Adds runtime roles like the default
// implicit role to RoleSet.
func FetchRoles(roleNames []string, access RoleGetter, traits map[string][]string) (RoleSet, error) {
return FetchRolesWithContext(roleNames, access, RoleTemplateContext{Traits: traits})
}
// FetchRolesWithContext fetches roles by their names, applies the role
// template context to role variables, and returns the RoleSet. Adds runtime
// roles like the default implicit role to RoleSet.
func FetchRolesWithContext(roleNames []string, access RoleGetter, ctx RoleTemplateContext) (RoleSet, error) {
roles, err := FetchRoleListWithContext(roleNames, access, ctx)
if err != nil {
return nil, trace.Wrap(err)
}
return NewRoleSet(roles...), nil
}
// FetchRolesForUser fetches a user's roles using their username and traits as
// role template context.
func FetchRolesForUser(user UserAccessState, access RoleGetter) (RoleSet, error) {
return FetchRolesWithContext(user.GetRoles(), access, RoleTemplateContext{
Username: user.GetName(),
Traits: user.GetTraits(),
})
}
// NewRoleSet returns new RoleSet based on the roles
func NewRoleSet(roles ...types.Role) RoleSet {
// unauthenticated Nop role should not have any privileges
// by default, otherwise it is too permissive
if len(roles) == 1 && roles[0].GetName() == string(types.RoleNop) {
return roles
}
return append(roles, NewImplicitRole())
}
// RoleSet is a set of roles that implements access control functionality
type RoleSet []types.Role
// EnumerationResult is a result of enumerating a role set against some property, e.g. allowed names or logins.
type EnumerationResult struct {
allowedDeniedMap map[string]bool
wildcardAllowed bool
wildcardDenied bool
}
func (result *EnumerationResult) filtered(value bool) []string {
var filtered []string
for entity, allow := range result.allowedDeniedMap {
if allow == value {
filtered = append(filtered, entity)
}
}
sort.Strings(filtered)
return filtered
}
// Denied returns all explicitly denied entities.
func (result *EnumerationResult) Denied() []string {
return result.filtered(false)
}
// Allowed returns all known allowed entities.
func (result *EnumerationResult) Allowed() []string {
if result.WildcardDenied() {
return nil
}
return result.filtered(true)
}
// WildcardAllowed is true if the * entity is allowed for a given rule set.
func (result *EnumerationResult) WildcardAllowed() bool {
return result.wildcardAllowed && !result.wildcardDenied
}
// WildcardDenied is true if the * entity is denied for a given rule set.
func (result *EnumerationResult) WildcardDenied() bool {
return result.wildcardDenied
}
// ToEntities converts result back to allowed and denied entity slices.
//
// If wildcard is denied, only "*" is returned for the denied slice.
// If wildcard is allowed, allowed entities will be appended to the allowed
// slice after the "*" as a hint for users to select.
// Denied entities is only included if the wildcard is allowed.
func (result *EnumerationResult) ToEntities() (allowed, denied []string) {
if result.wildcardDenied {
return nil, []string{types.Wildcard}
}
if result.wildcardAllowed {
return append([]string{types.Wildcard}, result.Allowed()...), result.Denied()
}
return result.Allowed(), nil
}
// NewEnumerationResult returns new EnumerationResult.
func NewEnumerationResult() EnumerationResult {
return EnumerationResult{
allowedDeniedMap: map[string]bool{},
wildcardAllowed: false,
wildcardDenied: false,
}
}
// NewEnumerationResultFromEntities creates a new EnumerationResult and
// populates the result with provided allowed and denied entries.
func NewEnumerationResultFromEntities(allowed, denied []string) EnumerationResult {
var wildcardAllowed bool
var wildcardDenied bool
allowedDeniedMap := make(map[string]bool)
for _, allow := range allowed {
if allow == types.Wildcard {
wildcardAllowed = true
} else {
allowedDeniedMap[allow] = true
}
}
for _, deny := range denied {
if deny == types.Wildcard {
wildcardDenied = true
wildcardAllowed = false
break
}
allowedDeniedMap[deny] = false
}
return EnumerationResult{
allowedDeniedMap: allowedDeniedMap,
wildcardAllowed: wildcardAllowed,
wildcardDenied: wildcardDenied,
}
}
// MatchNamespace returns true if given list of namespace matches
// target namespace, wildcard matches everything.
func MatchNamespace(selectors []string, namespace string) (bool, string) {
for _, n := range selectors {
if n == namespace || n == types.Wildcard {
return true, "matched"
}
}
return false, fmt.Sprintf("no match, role selectors %v, server namespace: %v", selectors, namespace)
}
// MatchAWSRoleARN returns true if provided role ARN matches selectors.
func MatchAWSRoleARN(selectors []string, roleARN string) (bool, string) {
if slices.Contains(selectors, roleARN) {
return true, "matched"
}
return false, fmt.Sprintf("no match, role selectors %v, role ARN: %v", selectors, roleARN)
}
// MatchAzureIdentity returns true if provided Azure identity matches selectors.
func MatchAzureIdentity(selectors []string, identity string, matchWildcard bool) (bool, string) {
identity = strings.ToLower(identity)
for _, l := range selectors {
if strings.ToLower(l) == identity {
return true, "element matched"
}
if matchWildcard && l == types.Wildcard {
return true, "wildcard matched"
}
}
return false, fmt.Sprintf("no match, role selectors %v, identity: %v", selectors, identity)
}
// MatchGCPServiceAccount returns true if provided GCP service account matches selectors.
func MatchGCPServiceAccount(selectors []string, account string, matchWildcard bool) (bool, string) {
for _, l := range selectors {
if l == account {
return true, "element matched"
}
if matchWildcard && l == types.Wildcard {
return true, "wildcard matched"
}
}
return false, fmt.Sprintf("no match, role selectors %v, identity: %v", selectors, account)
}
// MatchDatabaseName returns true if provided database name matches selectors.
func MatchDatabaseName(selectors []string, name string) (bool, string) {
for _, n := range selectors {
if n == name || n == types.Wildcard {
return true, "matched"
}
}
return false, fmt.Sprintf("no match, role selectors %v, database name: %v", selectors, name)
}
// MatchDatabaseUser returns true if provided database user matches selectors.
func MatchDatabaseUser(selectors []string, user string, matchWildcard, caseFold bool) (bool, string) {
for _, u := range selectors {
if caseFold {
if strings.EqualFold(u, user) {
return true, "matched"
}
} else if u == user {
return true, "matched"
}
if matchWildcard && u == types.Wildcard {
return true, "matched"
}
}
return false, fmt.Sprintf("no match, role selectors %v, database user: %v", selectors, user)
}
// MatchLabels matches selector against target. Empty selector matches
// nothing, wildcard matches everything.
func MatchLabels(selector types.Labels, target map[string]string) (bool, string, error) {
return MatchLabelGetter(selector, label.MapLabelGetter(target))
}
// MatchLabelGetter matches selector against labelGetter. Empty selector matches
// nothing, wildcard matches everything.
//
// Keep in sync with front-end implementation;
// - web/packages/teleport/src/Bots/Add/Shared/kubernetes.ts:34
func MatchLabelGetter(selector types.Labels, labelGetter label.LabelGetter) (bool, string, error) {
// Empty selector matches nothing.
if len(selector) == 0 {
return false, "no match, empty selector", nil
}
// *: * matches everything even empty target set.
selectorValues := selector[types.Wildcard]
if len(selectorValues) == 1 && selectorValues[0] == types.Wildcard {
return true, "matched", nil
}
// Perform full match.
for key, selectorValues := range selector {
targetVal, hasKey := labelGetter.GetLabel(key)
if !hasKey {
return false, fmt.Sprintf("no key match: '%v'", key), nil
}
if slices.Contains(selectorValues, types.Wildcard) {
continue
}
result, err := utils.SliceMatchesRegex(targetVal, selectorValues)
if err != nil {
return false, "", trace.Wrap(err)
} else if !result {
return false, fmt.Sprintf("no value match: got '%v' want: '%v'", targetVal, selectorValues), nil
}
}
return true, "matched", nil
}
// RoleNames returns a slice with role names. Removes runtime roles like
// the default implicit role.
func (set RoleSet) RoleNames() []string {
out := make([]string, 0, len(set))
for _, r := range set {
if r.GetName() == constants.DefaultImplicitRole {
continue
}
out = append(out, r.GetName())
}
return out
}
// Roles returns the list underlying roles this RoleSet is based on.
func (set RoleSet) Roles() []types.Role {
return slices.Clone(set)
}
// HasRole checks if the role set has the role
func (set RoleSet) HasRole(role string) bool {
for _, r := range set {
if r.GetName() == role {
return true
}
}
return false
}
// WithoutImplicit returns this role set with default implicit role filtered out.
func (set RoleSet) WithoutImplicit() (out RoleSet) {
for _, r := range set {
if r.GetName() == constants.DefaultImplicitRole {
continue
}
out = append(out, r)
}
return out
}
// PinSourceIP determines if the role set should use source IP pinning.
// If one or more roles in the set requires IP pinning then it will be enabled.
func (set RoleSet) PinSourceIP() bool {
for _, role := range set {
if role.GetOptions().PinSourceIP {
return true
}
}
return false
}
// GetAccessState returns the AccessState, setting [AccessState.MFARequired]
// according to the user's roles and cluster auth preference.
func (set RoleSet) GetAccessState(authPref readonly.AuthPreference) AccessState {
return AccessState{
MFARequired: set.getMFARequired(authPref.GetRequireMFAType()),
// We don't set EnableDeviceVerification here, as both it and DeviceVerified
// should be set in tandem.
}
}
func (set RoleSet) getMFARequired(clusterRequireMFAType types.RequireMFAType) MFARequired {
// MFA is always required according to the cluster auth pref.
if clusterRequireMFAType.IsSessionMFARequired() {
return MFARequiredAlways
}
// If MFA requirement is the same across all roles, we can skip the per-role check.
// Set mfaRequired to the first role's requirement, then check if all other roles match.
if len(set) > 0 {
rolesMFARequired := set[0].GetOptions().RequireMFAType.IsSessionMFARequired()
for _, role := range set[1:] {
if role.GetOptions().RequireMFAType.IsSessionMFARequired() != rolesMFARequired {
// This role differs from the MFA requirement of the other roles, return per-role.
return MFARequiredPerRole
}
}
if rolesMFARequired {
return MFARequiredAlways
}
}
// No roles to check or no roles require MFA.
return MFARequiredNever
}
// PrivateKeyPolicy returns the enforced private key policy for this role set.
func (set RoleSet) PrivateKeyPolicy(authPreferencePolicy keys.PrivateKeyPolicy) (keys.PrivateKeyPolicy, error) {
policySet := []keys.PrivateKeyPolicy{authPreferencePolicy}
for _, role := range set {
policySet = append(policySet, role.GetPrivateKeyPolicy())
}
return keys.PolicyThatSatisfiesSet(policySet)
}
// AdjustSessionTTL will reduce the requested ttl to the lowest max allowed TTL
// for this role set, otherwise it returns ttl unchanged
func (set RoleSet) AdjustSessionTTL(ttl time.Duration) time.Duration {
for _, role := range set {
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if maxSessionTTL != 0 && ttl > maxSessionTTL {
ttl = maxSessionTTL
}
}
return ttl
}
// AdjustMFAVerificationInterval will reduce the requested ttl to the lowest mfa verification interval
// for this role set if the role forces MFA tap, otherwise it returns ttl unchanged
func (set RoleSet) AdjustMFAVerificationInterval(ttl time.Duration, enforce bool) time.Duration {
for _, role := range set {
mfaVerificationInterval := role.GetOptions().MFAVerificationInterval
if role.GetOptions().RequireMFAType == types.RequireMFAType_OFF && !enforce {
continue
}
if mfaVerificationInterval != 0 && ttl > mfaVerificationInterval {
ttl = mfaVerificationInterval
}
}
return ttl
}
// MaxConnections returns the maximum number of concurrent ssh connections
// allowed. If MaxConnections is zero then no maximum was defined
// and the number of concurrent connections is unconstrained.
func (set RoleSet) MaxConnections() int64 {
var mcs int64
for _, role := range set {
if m := role.GetOptions().MaxConnections; m != 0 && (m < mcs || mcs == 0) {
mcs = m
}
}
return mcs
}
// MaxSessions returns the maximum number of concurrent ssh sessions
// per connection. If MaxSessions is zero then no maximum was defined
// and the number of sessions is unconstrained.
func (set RoleSet) MaxSessions() int64 {
var ms int64
for _, role := range set {
if m := role.GetOptions().MaxSessions; m != 0 && (m < ms || ms == 0) {
ms = m
}
}
return ms
}
// MaxKubernetesConnections implements [AccessChecker].
func (set RoleSet) MaxKubernetesConnections() int64 {
var mcs int64
for _, role := range set {
if m := role.GetOptions().MaxKubernetesConnections; m != 0 && (m < mcs || mcs == 0) {
mcs = m
}
}
return mcs
}
// SessionPolicySets returns the list of SessionPolicySets for all roles.
func (set RoleSet) SessionPolicySets() []*types.SessionTrackerPolicySet {
var policySets []*types.SessionTrackerPolicySet
for _, role := range set {
policySet := role.GetSessionPolicySet()
policySets = append(policySets, &policySet)
}
return policySets
}
// AdjustClientIdleTimeout adjusts requested idle timeout
// to the lowest max allowed timeout, the most restrictive
// option will be picked, negative values will be assumed as 0
func (set RoleSet) AdjustClientIdleTimeout(timeout time.Duration) time.Duration {
if timeout < 0 {
timeout = 0
}
for _, role := range set {
roleTimeout := role.GetOptions().ClientIdleTimeout
// 0 means not set, so it can't be most restrictive, disregard it too
if roleTimeout.Duration() <= 0 {
continue
}
switch {
// in case if timeout is 0, means that incoming value
// does not restrict the idle timeout, pick any other value
// set by the role
case timeout == 0:
timeout = roleTimeout.Duration()
case roleTimeout.Duration() < timeout:
timeout = roleTimeout.Duration()
}
}
return timeout
}
// AdjustDisconnectExpiredCert adjusts the value based on the role set
// the most restrictive option will be picked
func (set RoleSet) AdjustDisconnectExpiredCert(disconnect bool) bool {
for _, role := range set {
if role.GetOptions().DisconnectExpiredCert.Value() {
disconnect = true
}
}
return disconnect
}
// CheckKubeGroupsAndUsers check if role can login into kubernetes
// and returns two lists of allowed groups and users
func (set RoleSet) CheckKubeGroupsAndUsers(ttl time.Duration, overrideTTL bool, matchers ...RoleMatcher) ([]string, []string, error) {
groups := setutils.New[string]()
users := setutils.New[string]()
var matchedTTL bool
for _, role := range set {
ok, err := RoleMatchers(matchers).MatchAll(role, types.Allow)
if err != nil {
return nil, nil, trace.Wrap(err)
}
if !ok {
continue
}
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if overrideTTL || (ttl <= maxSessionTTL && maxSessionTTL != 0) {
matchedTTL = true
for _, group := range role.GetKubeGroups(types.Allow) {
groups.Add(group)
}
for _, user := range role.GetKubeUsers(types.Allow) {
users.Add(user)
}
}
}
for _, role := range set {
ok, _, err := RoleMatchers(matchers).MatchAny(role, types.Deny)
if err != nil {
return nil, nil, trace.Wrap(err)
}
if !ok {
continue
}
for _, group := range role.GetKubeGroups(types.Deny) {
groups.Remove(group)
}
for _, user := range role.GetKubeUsers(types.Deny) {
users.Remove(user)
}
}
if !matchedTTL {
return nil, nil, trace.AccessDenied("this user cannot request kubernetes access for %v", ttl)
}
if len(groups) == 0 && len(users) == 0 {
return nil, nil, trace.NotFound("this user cannot request kubernetes access, has no assigned groups or users")
}
return groups.Elements(), users.Elements(), nil
}
// CheckDatabaseNamesAndUsers checks if the role has any allowed database
// names or users.
func (set RoleSet) CheckDatabaseNamesAndUsers(ttl time.Duration, overrideTTL bool) ([]string, []string, error) {
names := setutils.New[string]()
users := setutils.New[string]()
var matchedTTL bool
for _, role := range set {
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if overrideTTL || (ttl <= maxSessionTTL && maxSessionTTL != 0) {
matchedTTL = true
for _, name := range role.GetDatabaseNames(types.Allow) {
names.Add(name)
}
for _, user := range role.GetDatabaseUsers(types.Allow) {
users.Add(user)
}
}
}
for _, role := range set {
for _, name := range role.GetDatabaseNames(types.Deny) {
names.Remove(name)
}
for _, user := range role.GetDatabaseUsers(types.Deny) {
users.Remove(user)
}
}
if !matchedTTL {
return nil, nil, trace.AccessDenied("this user cannot request database access for %v", ttl)
}
if len(names) == 0 && len(users) == 0 {
return nil, nil, trace.NotFound("this user cannot request database access, has no assigned database names or users")
}
return names.Elements(), users.Elements(), nil
}
// CheckAWSRoleARNs returns a list of AWS role ARNs this role set is allowed to assume.
func (set RoleSet) CheckAWSRoleARNs(ttl time.Duration, overrideTTL bool) ([]string, error) {
arns := setutils.New[string]()
var matchedTTL bool
for _, role := range set {
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if overrideTTL || (ttl <= maxSessionTTL && maxSessionTTL != 0) {
matchedTTL = true
for _, arn := range role.GetAWSRoleARNs(types.Allow) {
arns.Add(arn)
}
}
}
for _, role := range set {
for _, arn := range role.GetAWSRoleARNs(types.Deny) {
arns.Remove(arn)
}
}
if !matchedTTL {
return nil, trace.AccessDenied("this user cannot request AWS management console access for %v", ttl)
}
if len(arns) == 0 {
return nil, trace.NotFound("this user cannot request AWS management console, has no assigned role ARNs")
}
return arns.Elements(), nil
}
// CheckAzureIdentities returns a list of Azure identities the user is allowed to assume.
func (set RoleSet) CheckAzureIdentities(ttl time.Duration, overrideTTL bool) ([]string, error) {
identities := make(map[string]string)
var matchedTTL bool
for _, role := range set {
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if overrideTTL || (ttl <= maxSessionTTL && maxSessionTTL != 0) {
matchedTTL = true
for _, identity := range role.GetAzureIdentities(types.Allow) {
identities[strings.ToLower(identity)] = identity
}
}
}
for _, role := range set {
for _, identity := range role.GetAzureIdentities(types.Deny) {
// deny * cleans options
if identity == types.Wildcard {
identities = make(map[string]string)
}
// remove particular identity
delete(identities, strings.ToLower(identity))
}
}
if !matchedTTL {
return nil, trace.AccessDenied("this user cannot access Azure API for %v", ttl)
}
if len(identities) == 0 {
return nil, trace.NotFound("this user cannot access Azure API, has no assigned identities")
}
out := make([]string, 0, len(identities))
for _, identity := range identities {
out = append(out, identity)
}
sort.Strings(out)
return out, nil
}
// CheckGCPServiceAccounts returns a list of GCP service accounts this role set is allowed to assume.
func (set RoleSet) CheckGCPServiceAccounts(ttl time.Duration, overrideTTL bool) ([]string, error) {
accounts := setutils.New[string]()
var matchedTTL bool
for _, role := range set {
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if overrideTTL || (ttl <= maxSessionTTL && maxSessionTTL != 0) {
matchedTTL = true
for _, account := range role.GetGCPServiceAccounts(types.Allow) {
accounts.Add(strings.ToLower(account))
}
}
}
for _, role := range set {
for _, account := range role.GetGCPServiceAccounts(types.Deny) {
// deny * removes all accounts
if account == types.Wildcard {
accounts = setutils.New[string]()
}
// remove particular account
accounts.Remove(strings.ToLower(account))
}
}
if !matchedTTL {
return nil, trace.AccessDenied("this user cannot request GCP API access for %v", ttl)
}
if len(accounts) == 0 {
return nil, trace.NotFound("this user cannot request GCP API access, has no assigned service accounts")
}
return accounts.Elements(), nil
}
// checkAccessToSAMLIdPLegacy checks access to the SAML IdP based on
// IDP enabled/disabled in role option and MFA. The IDP option is enforced
// in Teleport role version v7 and below.
func checkAccessToSAMLIdPLegacy(state AccessState, role types.Role) error {
ctx := context.Background()
if state.MFARequired == MFARequiredAlways && !state.MFAVerified {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access to SAML IdP denied, cluster requires per-session MFA")
return trace.Wrap(ErrSessionMFARequired)
}
mfaAllowed := state.MFAVerified || state.MFARequired == MFARequiredNever
options := role.GetOptions()
// This should never happen, but we should make sure that we don't get a nil pointer error here.
if options.IDP == nil || options.IDP.SAML == nil || options.IDP.SAML.Enabled == nil {
return nil
}
// If any role specifically denies access to the IdP, we'll return AccessDenied.
if !options.IDP.SAML.Enabled.Value {
return trace.AccessDenied("user has been denied access to the SAML IdP by role %s", role.GetName())
}
if !mfaAllowed && options.RequireMFAType.IsSessionMFARequired() {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access to SAML IdP denied, role requires per-session MFA",
slog.String("role", role.GetName()),
)
return trace.Wrap(ErrSessionMFARequired)
}
return nil
}
// CheckAccessToSAMLIdP checks access to SAML service provider resource.
// For Teleport role version v7 and below (legacy SAML IdP RBAC), only MFA
// and IDP role option is checked.
// For Teleport role version v8 and above (non-legacy SAML IdP RBAC),
// labels, MFA and Device Trust is checked.
// IDP option in the auth preference is checked in both the cases.
func (set RoleSet) CheckAccessToSAMLIdP(r AccessCheckable, username string, traits wrappers.Traits, authPref readonly.AuthPreference, state AccessState, matchers ...RoleMatcher) error {
if authPref != nil {
if !authPref.IsSAMLIdPEnabled() {
return trace.AccessDenied("SAML IdP is disabled at the cluster level")
}
}
if len(set) == 0 {
return trace.AccessDenied("access to %v denied. User does not have permissions. %v",
r.GetKind(), "No roles assigned to user")
}
var v8RoleSet RoleSet
for _, role := range set {
if !types.IsLegacySAMLRBAC(role.GetVersion()) {
v8RoleSet = append(v8RoleSet, role)
continue
}
if err := checkAccessToSAMLIdPLegacy(state, role); err != nil {
return trace.Wrap(err)
}
}
// We checked for empty roleset early on this method. Reaching this part
// and zero non-legacy roleset means that the user was allowed access
// with legacy roles. We'll honor that and return, otherwise, checkAccess
// will deny access on empty role set.
if len(v8RoleSet) == 0 {
return nil
}
if _, err := v8RoleSet.checkAccess(r, username, traits, state, matchers...); err != nil {
return trace.Wrap(err)
}
return nil
}
// CheckLoginDuration checks if role set can login up to given duration and
// returns a combined list of allowed logins.
func (set RoleSet) CheckLoginDuration(ttl time.Duration) ([]string, error) {
logins, matchedTTL := set.GetLoginsForTTL(ttl)
if !matchedTTL {
return nil, trace.AccessDenied("this user cannot request a certificate for %v", ttl)
}
if len(logins) == 0 && !set.hasPossibleLogins() {
// user was deliberately configured to have no login capability,
// but ssh certificates must contain at least one valid principal.
// we add a single distinctive value which should be unique, and
// will never be a valid unix login (due to leading '-').
logins = []string{constants.NoLoginPrefix + uuid.New().String()}
}
if len(logins) == 0 {
return nil, trace.AccessDenied("this user cannot create SSH sessions, has no allowed logins")
}
return logins, nil
}
// GetAllLogins returns all valid unix logins for the RoleSet.
func (set RoleSet) GetAllLogins() []string {
logins, _ := set.GetLoginsForTTL(0)
return logins
}
// GetLoginsForTTL collects all logins that are valid for the given TTL. The matchedTTL
// value indicates whether the TTL is within scope of *any* role. This helps to distinguish
// between TTLs which are categorically invalid, and TTLs which are theoretically valid
// but happen to grant no logins.
func (set RoleSet) GetLoginsForTTL(ttl time.Duration) (logins []string, matchedTTL bool) {
for _, role := range set {
maxSessionTTL := role.GetOptions().MaxSessionTTL.Value()
if ttl <= maxSessionTTL && maxSessionTTL != 0 {
matchedTTL = true
logins = append(logins, role.GetLogins(types.Allow)...)
}
}
return apiutils.Deduplicate(logins), matchedTTL
}
func (set RoleSet) hasPossibleLogins() bool {
for _, role := range set {
if role.GetName() == constants.DefaultImplicitRole {
continue
}
if len(role.GetLogins(types.Allow)) != 0 {
return true
}
}
return false
}
// AWSRoleARNMatcher matches a role against AWS role ARN.
type AWSRoleARNMatcher struct {
RoleARN string
}
// Match matches AWS role ARN against provided role and condition.
func (m *AWSRoleARNMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
match, _ := MatchAWSRoleARN(role.GetAWSRoleARNs(condition), m.RoleARN)
return match, nil
}
// String returns the matcher's string representation.
func (m *AWSRoleARNMatcher) String() string {
return fmt.Sprintf("AWSRoleARNMatcher(RoleARN=%v)", m.RoleARN)
}
// AzureIdentityMatcher matches a role against Azure identity.
type AzureIdentityMatcher struct {
Identity string
}
// Match matches Azure identity against provided role and condition.
func (m *AzureIdentityMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
match, _ := MatchAzureIdentity(role.GetAzureIdentities(condition), m.Identity, condition == types.Deny)
return match, nil
}
// String returns the matcher's string representation.
func (m *AzureIdentityMatcher) String() string {
return fmt.Sprintf("AzureIdentityMatcher(Identity=%v)", m.Identity)
}
// GCPServiceAccountMatcher matches a role against GCP service account.
type GCPServiceAccountMatcher struct {
// ServiceAccount is a GCP service account to match, e.g. teleport@example-123456.iam.gserviceaccount.com.
// It can also be a wildcard *, but that is only respected for Deny rules.
ServiceAccount string
}
// Match matches GCP ServiceAccount against provided role and condition.
func (m *GCPServiceAccountMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
match, _ := MatchGCPServiceAccount(role.GetGCPServiceAccounts(condition), m.ServiceAccount, condition == types.Deny)
return match, nil
}
// String returns the matcher's string representation.
func (m *GCPServiceAccountMatcher) String() string {
return fmt.Sprintf("GCPServiceAccountMatcher(ServiceAccount=%v)", m.ServiceAccount)
}
// CanImpersonateSomeone returns true if this checker has any impersonation rules
func (set RoleSet) CanImpersonateSomeone() bool {
for _, role := range set {
cond := role.GetImpersonateConditions(types.Allow)
if !cond.IsEmpty() {
return true
}
}
return false
}
// CheckImpersonate returns nil if this role set can impersonate
// a user and their roles, returns AccessDenied otherwise
// CheckImpersonate checks whether current user is allowed to impersonate
// users and roles
func (set RoleSet) CheckImpersonate(currentUser, impersonateUser types.User, impersonateRoles []types.Role) error {
ctx := &impersonateContext{
user: currentUser,
impersonateUser: impersonateUser,
}
whereParser, err := newImpersonateWhereParser(ctx)
if err != nil {
return trace.Wrap(err)
}
// check deny: a single match on a deny rule prohibits access
for _, role := range set {
cond := role.GetImpersonateConditions(types.Deny)
matched, err := matchDenyImpersonateCondition(cond, impersonateUser, impersonateRoles)
if err != nil {
return trace.Wrap(err)
}
if matched {
return trace.AccessDenied("access denied to '%s' to impersonate user '%s' and roles '%s'", currentUser.GetName(), impersonateUser.GetName(), roleNames(impersonateRoles))
}
}
// check allow: if matches, allow to impersonate
for _, role := range set {
cond := role.GetImpersonateConditions(types.Allow)
matched, err := matchAllowImpersonateCondition(ctx, whereParser, cond, impersonateUser, impersonateRoles)
if err != nil {
return trace.Wrap(err)
}
if matched {
return nil
}
}
return trace.AccessDenied("access denied to '%s' to impersonate user '%s' and roles '%s'", currentUser.GetName(), impersonateUser.GetName(), roleNames(impersonateRoles))
}
// CheckImpersonateRoles validates that the current user can perform role-only impersonation
// of the given roles. Role-only impersonation requires an allow rule with
// roles but no users (and no user-less deny rules). All requested roles must
// be allowed for the check to succeed.
func (set RoleSet) CheckImpersonateRoles(currentUser types.User, impersonateRoles []types.Role) error {
ctx := &impersonateContext{
user: currentUser,
}
whereParser, err := newImpersonateWhereParser(ctx)
if err != nil {
return trace.Wrap(err)
}
// TODO: Unlike regular impersonation where all requested roles must be
// granted by a single impersonation role, it would be reasonable to
// request several roles whose `allow` conditions are split between
// several roles. Our initial use-case doesn't require this, so for now
// we'll assume all requested roles must be granted by a single `allow`.
// check deny: a single match on a deny rule prohibits access
for _, role := range set {
matched, err := matchDenyRoleImpersonateCondition(role, impersonateRoles)
if err != nil {
return trace.Wrap(err)
}
if matched {
return trace.AccessDenied("access denied to '%s' to impersonate roles '%s'", currentUser.GetName(), roleNames(impersonateRoles))
}
}
// check allow: if any one Role satisfies all the role requests, allow impersonation
for _, role := range set {
matched, err := matchAllowRoleImpersonateCondition(ctx, whereParser, role, impersonateRoles)
if err != nil {
return trace.Wrap(err)
}
if matched {
return nil
}
}
return trace.AccessDenied("access denied to '%s' to impersonate roles '%s'", currentUser.GetName(), roleNames(impersonateRoles))
}
// CheckSubmitForUser checks whether the current user is allowed to
// submit reviews for other users, to be used by plugins.
func (set RoleSet) CheckSubmitForUser(currentUser, submitForUser types.User) error {
for _, role := range set {
denyUsers := role.GetSubmitForUsers(types.Deny)
anyDenyUser, err := parse.NewAnyMatcher(denyUsers)
if err != nil {
return trace.Wrap(err)
}
if anyDenyUser.Match(submitForUser.GetName()) {
return trace.AccessDenied("access denied for '%s' to submit for user '%s'", currentUser.GetName(), submitForUser.GetName())
}
}
for _, role := range set {
allowUsers := role.GetSubmitForUsers(types.Allow)
anyAllowUser, err := parse.NewAnyMatcher(allowUsers)
if err != nil {
return trace.Wrap(err)
}
if anyAllowUser.Match(submitForUser.GetName()) {
return nil
}
}
return trace.AccessDenied("access denied for '%s' to submit for user '%s'", currentUser.GetName(), submitForUser.GetName())
}
// LockingMode returns the locking mode to apply with this RoleSet.
func (set RoleSet) LockingMode(defaultMode constants.LockingMode) constants.LockingMode {
mode := defaultMode
for _, role := range set {
options := role.GetOptions()
if options.Lock == constants.LockingModeStrict {
return constants.LockingModeStrict
}
if options.Lock != "" {
mode = options.Lock
}
}
return mode
}
// CertificateExtensions returns the list of extensions for each role in the RoleSet
func (set RoleSet) CertificateExtensions() []*types.CertExtension {
var exts []*types.CertExtension
for _, role := range set {
exts = append(exts, role.GetOptions().CertExtensions...)
}
return exts
}
// SessionRecordingMode returns the recording mode for a specific service.
func (set RoleSet) SessionRecordingMode(service constants.SessionRecordingService) constants.SessionRecordingMode {
defaultValue := constants.SessionRecordingModeBestEffort
useDefault := true
for _, role := range set {
recordSession := role.GetOptions().RecordSession
// If one of the default values is "strict", set it as the value.
if recordSession.Default == constants.SessionRecordingModeStrict {
defaultValue = constants.SessionRecordingModeStrict
}
var roleMode constants.SessionRecordingMode
switch service {
case constants.SessionRecordingServiceSSH:
roleMode = recordSession.SSH
}
switch roleMode {
case constants.SessionRecordingModeStrict:
// Early return as "strict" since it is the strictest value.
return constants.SessionRecordingModeStrict
case constants.SessionRecordingModeBestEffort:
useDefault = false
}
}
// Return the strictest default value.
if useDefault {
return defaultValue
}
return constants.SessionRecordingModeBestEffort
}
func contains[S ~[]E, E any](s S, f func(E) (bool, error)) (bool, error) {
for i := range s {
match, err := f(s[i])
if err != nil {
return false, err
}
if match {
return true, nil
}
}
return false, nil
}
// matchSPIFFESVIDDenyConditions compares a slice of SPIFFE Role Conditions against
// a requested SPIFFE SVID generation. Any field within a condition must match,
// and any condition in the slice can match for the function to return true.
func matchSPIFFESVIDDenyConditions(
conds []*types.SPIFFERoleCondition,
spiffeIDPath string,
dnsSANs []string,
ipSANs []net.IP,
) (bool, error) {
return contains(conds, func(cond *types.SPIFFERoleCondition) (bool, error) {
// Match SPIFFE ID path.
isPathMatch, err := utils.MatchString(spiffeIDPath, cond.Path)
if err != nil {
return false, trace.Wrap(err)
}
if !isPathMatch {
return false, nil
}
// if any DNS SAN in the condition matches, we say the DNS SAN part
// of the condition matches
isDNSMatch := true
for _, dnsSANMatcher := range cond.DNSSANs {
isDNSMatch, err = contains(dnsSANs, func(reqDNSSAN string) (bool, error) {
return utils.MatchString(reqDNSSAN, dnsSANMatcher)
})
if err != nil {
return false, trace.Wrap(err)
}
if isDNSMatch {
break
}
}
if !isDNSMatch {
return false, nil
}
// Any IP SAN requested can match one of the IP SAN matchers in the
// condition.
isIPMatch := true
for _, ipSANMatcher := range cond.IPSANs {
isIPMatch, err = contains(ipSANs, func(reqIPSAN net.IP) (bool, error) {
_, cidr, err := net.ParseCIDR(ipSANMatcher)
if err != nil {
return false, trace.Wrap(err, "parsing cidr")
}
return cidr.Contains(reqIPSAN), nil
})
if err != nil {
return false, trace.Wrap(err)
}
if isIPMatch {
break
}
}
// all other conditions met
return isIPMatch, nil
})
}
// matchSPIFFESVIDAllowConditions compares a slice of SPIFFE Role Conditions against
// a requested SPIFFE SVID generation. All fields within a condition must match,
// but any condition in the slice can match for the function to return true.
func matchSPIFFESVIDAllowConditions(
conds []*types.SPIFFERoleCondition,
spiffeIDPath string,
dnsSANs []string,
ipSANs []net.IP,
) (bool, error) {
return contains(conds, func(cond *types.SPIFFERoleCondition) (bool, error) {
// Match SPIFFE ID path.
match, err := utils.MatchString(spiffeIDPath, cond.Path)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
// No match - skip to next condition.
return false, nil
}
// All DNS SANs requested must match one of the DNS SAN matchers in the
// condition.
for _, dnsSAN := range dnsSANs {
match, err := contains(cond.DNSSANs, func(s string) (bool, error) {
match, err := utils.MatchString(dnsSAN, s)
if err != nil {
return false, trace.Wrap(err)
}
return match, nil
})
if err != nil {
return false, trace.Wrap(err)
}
if !match {
return false, nil
}
}
// All IP SANs requested must match one of the IP SAN matchers in the
// condition.
for _, ipSAN := range ipSANs {
match, err := contains(cond.IPSANs, func(s string) (bool, error) {
_, cidr, err := net.ParseCIDR(s)
if err != nil {
return false, trace.Wrap(err, "parsing cidr")
}
return cidr.Contains(ipSAN), nil
})
if err != nil {
return false, trace.Wrap(err)
}
if !match {
return false, nil
}
}
// All condition fields matched.
return true, nil
})
}
// CheckSPIFFESVID checks if the role set has access to generating the
// requested SPIFFE ID. Returns an error if the role set does not have the
// ability to generate the requested SVID.
func (set RoleSet) CheckSPIFFESVID(spiffeIDPath string, dnsSANs []string, ipSANs []net.IP) error {
accessDenied := trace.AccessDenied("access denied to generate SVID %q", spiffeIDPath)
// check deny: a single match on a deny rule prohibits generation
for _, role := range set {
cond := role.GetSPIFFEConditions(types.Deny)
matched, err := matchSPIFFESVIDDenyConditions(cond, spiffeIDPath, dnsSANs, ipSANs)
if err != nil {
return trace.Wrap(err)
}
if matched {
return accessDenied
}
}
// check allow: if a single condition matches, allow generation
for _, role := range set {
cond := role.GetSPIFFEConditions(types.Allow)
matched, err := matchSPIFFESVIDAllowConditions(cond, spiffeIDPath, dnsSANs, ipSANs)
if err != nil {
return trace.Wrap(err)
}
if matched {
return nil
}
}
return accessDenied
}
func roleNames(roles []types.Role) string {
out := make([]string, len(roles))
for i := range roles {
out[i] = roles[i].GetName()
}
return strings.Join(out, ", ")
}
// matchAllowImpersonateCondition matches impersonate condition,
// both user, role and where condition has to match
func matchAllowImpersonateCondition(ctx *impersonateContext, whereParser predicate.Parser, cond types.ImpersonateConditions, impersonateUser types.User, impersonateRoles []types.Role) (bool, error) {
// User impersonation requires both users and roles. Roles with no users
// must use RoleRequests instead; however, we can't treat this as an error
// since this function is tested against all roles regardless of how
// they'll be used.
if len(cond.Users) == 0 || len(cond.Roles) == 0 {
return false, nil
}
anyUser, err := parse.NewAnyMatcher(cond.Users)
if err != nil {
return false, trace.Wrap(err)
}
if !anyUser.Match(impersonateUser.GetName()) {
return false, nil
}
anyRole, err := parse.NewAnyMatcher(cond.Roles)
if err != nil {
return false, trace.Wrap(err)
}
for _, impersonateRole := range impersonateRoles {
if !anyRole.Match(impersonateRole.GetName()) {
return false, nil
}
// TODO:
// This set impersonateRole inside the ctx that is in turn used inside whereParser
// which is created in CheckImpersonate above but is being used right below.
// This is unfortunate interface of the parser, instead
// parser should accept additional context as a first argument.
ctx.impersonateRole = impersonateRole
match, err := matchesImpersonateWhere(cond, whereParser)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
return false, nil
}
}
return true, nil
}
// matchDenyImpersonateCondition matches impersonate condition,
// greedy is used for deny type rules, where any user or role can match
func matchDenyImpersonateCondition(cond types.ImpersonateConditions, impersonateUser types.User, impersonateRoles []types.Role) (bool, error) {
// As above, user impersonation requires both users and roles. We can't
// return an error to ensure role impersonation rules are allowed to exist
// in the system.
if len(cond.Users) == 0 || len(cond.Roles) == 0 {
return false, nil
}
anyUser, err := parse.NewAnyMatcher(cond.Users)
if err != nil {
return false, trace.Wrap(err)
}
if anyUser.Match(impersonateUser.GetName()) {
return true, nil
}
anyRole, err := parse.NewAnyMatcher(cond.Roles)
if err != nil {
return false, trace.Wrap(err)
}
for _, impersonateRole := range impersonateRoles {
if anyRole.Match(impersonateRole.GetName()) {
return true, nil
}
}
return false, nil
}
// matchAllowRoleImpersonateCondition matches an allow impersonate condition
// specifically for role-only impersonation, where only roles are matched.
func matchAllowRoleImpersonateCondition(ctx *impersonateContext, whereParser predicate.Parser, role types.Role, impersonateRoles []types.Role) (bool, error) {
cond := role.GetImpersonateConditions(types.Allow)
// an empty set matches nothing
if len(cond.Users) == 0 && len(cond.Roles) == 0 {
return false, nil
}
// Role impersonation can never apply to users.
if len(cond.Users) != 0 {
slog.WarnContext(context.Background(),
"Allow rule did not match due to users being set. For role-only impersonation, only roles should be set in allow/deny rules.",
"role", role.GetName(),
)
return false, nil
}
// By this point, at least 1 role is guaranteed.
anyRole, err := parse.NewAnyMatcher(cond.Roles)
if err != nil {
return false, trace.Wrap(err)
}
for _, impersonateRole := range impersonateRoles {
if !anyRole.Match(impersonateRole.GetName()) {
return false, nil
}
// TODO:
// This set impersonateRole inside the ctx that is in turn used inside whereParser
// which is created in CheckImpersonate above but is being used right below.
// This is unfortunate interface of the parser, instead
// parser should accept additional context as a first argument.
ctx.impersonateRole = impersonateRole
match, err := matchesImpersonateWhere(cond, whereParser)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
return false, nil
}
}
return true, nil
}
// matchDenyRoleImpersonateCondition matches a deny impersonate condition
// specifically for role impersonation, where only roles are matched.
func matchDenyRoleImpersonateCondition(role types.Role, impersonateRoles []types.Role) (bool, error) {
cond := role.GetImpersonateConditions(types.Deny)
// an empty set matches nothing
if len(cond.Users) == 0 && len(cond.Roles) == 0 {
return false, nil
}
// If any users are defined in a role-impersonation deny rule, it always
// matches. This functionally disables role impersonation for rules
// containing a `users` deny entry, which is acceptable because only bots
// should ever use role impersonation.
if len(cond.Users) != 0 {
slog.WarnContext(context.Background(),
"Deny rule matched due to users being set. For role-only impersonation, only roles should be set in allow/deny rules.",
"role", role.GetName(),
)
return true, nil
}
// By this point, at least 1 role is guaranteed.
anyRole, err := parse.NewAnyMatcher(cond.Roles)
if err != nil {
return false, trace.Wrap(err)
}
for _, impersonateRole := range impersonateRoles {
if anyRole.Match(impersonateRole.GetName()) {
return true, nil
}
}
return false, nil
}
// RoleMatcherFunc is a convenience type for creating a role matcher from a function.
type RoleMatcherFunc func(types.Role, types.RoleConditionType) (bool, error)
func (f RoleMatcherFunc) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
return f(role, condition)
}
// RoleMatcher defines an interface for a generic role matcher.
type RoleMatcher interface {
Match(types.Role, types.RoleConditionType) (bool, error)
}
// RoleMatchers defines a list of matchers.
type RoleMatchers []RoleMatcher
// MatchAll returns true if all matchers in the set match.
func (m RoleMatchers) MatchAll(role types.Role, condition types.RoleConditionType) (bool, error) {
for _, matcher := range m {
match, err := matcher.Match(role, condition)
if err != nil {
return false, trace.Wrap(err)
}
if !match {
return false, nil
}
}
return true, nil
}
// MatchAny returns true if at least one of the matchers in the set matches.
//
// If the result is true, returns matcher that matched.
func (m RoleMatchers) MatchAny(role types.Role, condition types.RoleConditionType) (bool, RoleMatcher, error) {
for _, matcher := range m {
match, err := matcher.Match(role, condition)
if err != nil {
return false, nil, trace.Wrap(err)
}
if match {
return true, matcher, nil
}
}
return false, nil, nil
}
// AnyOf returns a RoleMatcher that succeeds if ANY of the underlying matchers match.
func (m RoleMatchers) AnyOf() RoleMatcher {
return RoleMatcherFunc(func(r types.Role, cond types.RoleConditionType) (bool, error) {
ok, _, err := m.MatchAny(r, cond)
return ok, err
})
}
// databaseUserMatcher matches a role against database account name.
type databaseUserMatcher struct {
// user is the name of the database user.
user string
// alternativeNames is a list of alternative names for the database user.
alternativeNames []string
// caseInsensitive specifies if the username is case insensitive.
caseInsensitive bool
}
// NewDatabaseUserMatcher creates a RoleMatcher that checks whether the role's
// database users match the specified condition.
func NewDatabaseUserMatcher(db types.Database, user string) RoleMatcher {
if db.SupportAWSIAMRoleARNAsUsers() {
return &databaseUserMatcher{
user: user,
alternativeNames: makeUsernamesForAWSRoleARN(db, user),
caseInsensitive: db.IsUsernameCaseInsensitive(),
}
}
if db.RequireAWSIAMRolesAsUsers() {
return &databaseUserMatcher{
user: user,
alternativeNames: makeAlternativeNamesForAWSRole(db, user),
caseInsensitive: db.IsUsernameCaseInsensitive(),
}
}
return &databaseUserMatcher{
user: user,
caseInsensitive: db.IsUsernameCaseInsensitive(),
}
}
// Match matches database account name against provided role and condition.
func (m *databaseUserMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
selectors := role.GetDatabaseUsers(condition)
if match, _ := MatchDatabaseUser(selectors, m.user, true /*matchWildcard*/, m.caseInsensitive); match {
return true, nil
}
for _, altName := range m.alternativeNames {
if match, _ := MatchDatabaseUser(selectors, altName, false /*matchWildcard*/, m.caseInsensitive); match {
return true, nil
}
}
return false, nil
}
// String returns the matcher's string representation.
func (m *databaseUserMatcher) String() string {
return fmt.Sprintf("databaseUserMatcher(user=%v, alternativeNames=%v)", m.user, m.alternativeNames)
}
func makeAlternativeNamesForAWSRole(db types.Database, user string) []string {
metadata := db.GetAWS()
if metadata.Region == "" || metadata.AccountID == "" {
return nil
}
// If input database user is a role ARN, try the short role name.
// The input role ARN must have matching partition and account ID in
// order to try the short role name.
if arn.IsARN(user) {
roleName, err := awsutils.ValidateRoleARNAndExtractRoleName(user, metadata.Partition(), metadata.AccountID)
if err != nil {
return nil
}
return []string{roleName}
}
// If input database user is the short role name, try the full ARN.
roleARN, err := awsutils.BuildRoleARN(user, metadata.Region, metadata.AccountID)
if err != nil {
return nil
}
return []string{roleARN}
}
// makeUsernamesForAWSRoleARN builds ARN alternatives for database users who are full or
// partial ARN.
func makeUsernamesForAWSRoleARN(db types.Database, user string) []string {
if !awsutils.IsRoleARN(user) {
return nil
}
metadata := db.GetAWS()
if metadata.Region != "" && metadata.AccountID != "" && awsutils.IsPartialRoleARN(user) {
roleARN, err := awsutils.BuildRoleARN(user, metadata.Region, metadata.AccountID)
if err != nil {
return nil
}
return []string{roleARN}
}
roleARN, err := awsutils.ParseRoleARN(user)
if err != nil {
return nil
}
return []string{roleARN.Resource}
}
// DatabaseNameMatcher matches a role against database name.
type DatabaseNameMatcher struct {
Name string
}
// Match matches database name against provided role and condition.
func (m *DatabaseNameMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
match, _ := MatchDatabaseName(role.GetDatabaseNames(condition), m.Name)
return match, nil
}
// String returns the matcher's string representation.
func (m *DatabaseNameMatcher) String() string {
return fmt.Sprintf("DatabaseNameMatcher(Name=%v)", m.Name)
}
type loginMatcher struct {
login string
}
// NewLoginMatcher creates a RoleMatcher that checks whether the role's logins
// match the specified condition.
func NewLoginMatcher(login string) RoleMatcher {
return &loginMatcher{login: login}
}
// Match matches a login against a role.
func (l *loginMatcher) Match(role types.Role, typ types.RoleConditionType) (bool, error) {
logins := role.GetLogins(typ)
if slices.Contains(logins, l.login) {
return true, nil
}
return false, nil
}
type windowsLoginMatcher struct {
login string
}
// NewWindowsLoginMatcher creates a RoleMatcher that checks whether the role's
// Windows desktop logins match the specified condition.
func NewWindowsLoginMatcher(login string) RoleMatcher {
return &windowsLoginMatcher{login: login}
}
// Match matches a Windows Desktop login against a role.
func (l *windowsLoginMatcher) Match(role types.Role, typ types.RoleConditionType) (bool, error) {
logins := role.GetWindowsLogins(typ)
if slices.Contains(logins, l.login) {
return true, nil
}
return false, nil
}
type linuxDesktopLoginMatcher struct {
login string
}
// NewLinuxDesktopLoginMatcher creates a RoleMatcher that checks whether the role's
// Linux desktop logins match the specified condition.
func NewLinuxDesktopLoginMatcher(login string) RoleMatcher {
return &linuxDesktopLoginMatcher{login: login}
}
// Match matches a Linux Desktop login against a role.
func (l *linuxDesktopLoginMatcher) Match(role types.Role, typ types.RoleConditionType) (bool, error) {
logins := role.GetLinuxDesktopLogins(typ)
return slices.Contains(logins, l.login), nil
}
type awsAppLoginMatcher struct {
awsRole string
}
// NewAppAWSLoginMatcher creates a RoleMatcher that checks whether the role's
// AWS Role ARN match the specified condition.
func NewAppAWSLoginMatcher(awsRole string) RoleMatcher {
return &awsAppLoginMatcher{awsRole: awsRole}
}
// Match matches an AWS Role ARN login against a role.
func (l *awsAppLoginMatcher) Match(role types.Role, typ types.RoleConditionType) (bool, error) {
awsRoles := role.GetAWSRoleARNs(typ)
if slices.Contains(awsRoles, l.awsRole) {
return true, nil
}
return false, nil
}
type kubernetesClusterLabelMatcher struct {
clusterLabels map[string]string
username string
userTraits wrappers.Traits
}
// NewKubeResourcesMatcher creates a new KubeResourcesMatcher matcher that
// matches a role against any Kubernetes Resource specified.
// It also keeps track of the resources that did not match any of user's roles and
// that shouldn't be included in the resource ids because the user is not allowed
// to request them.
func NewKubeResourcesMatcher(resources []types.KubernetesResource) *KubeResourcesMatcher {
matcher := &KubeResourcesMatcher{
resources: resources,
unmatchedReqs: map[string]struct{}{},
}
for _, r := range resources {
matcher.unmatchedReqs[unmatchedKey(r)] = struct{}{}
}
return matcher
}
// unmatchedKey returns a unique key for a Kubernetes resource.
// It is used to keep track of the resources that did not match any of user's roles.
// Format: <kind>/<namespace>/<name>
func unmatchedKey(r types.KubernetesResource) string {
return path.Join(r.Kind, r.ClusterResource())
}
// KubeResourcesMatcher matches a role against any Kubernetes Resource specified.
// It also keeps track of the resources that did not match any of user's roles and
// that shouldn't be included in the resource ids because the user is not allowed
// to request them.
type KubeResourcesMatcher struct {
resources []types.KubernetesResource
unmatchedReqs map[string]struct{}
}
// Match matches a Kubernetes resource against provided role and condition.
func (m *KubeResourcesMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
var finalResult bool
for _, resource := range m.resources {
// We use utils.KubeResourceMatchesRegexWithVerbsCollector instead of utils.KubeResourceMatchesRegex
// because KubeResourcesMatcher is used to match access request resources at creation time against
// the roles specified in the `search_as_roles` field. This means that we don't have the request verb
// at this point and we need to match the resource against all the verbs specified in the role.
// If the resource matches any of the verbs, we consider the resource as matched.
// Verbs are enforced at the request time when the user is trying to access the Kubernetes Pod.
result, _, err := utils.KubeResourceMatchesRegexWithVerbsCollector(resource, role.GetKubeResources(condition))
if err != nil {
return false, trace.Wrap(err)
}
if result {
delete(m.unmatchedReqs, unmatchedKey(resource))
finalResult = true
}
}
return finalResult, nil
}
// String returns the matcher's string representation.
func (m *KubeResourcesMatcher) String() string {
return fmt.Sprintf("KubeResourcesMatcher(Resources=%v)", m.resources)
}
// Unmatched returns the Kubernetes Resource request access that that didn't
// match with any `search_as_roles` kubernetes resources.
func (m *KubeResourcesMatcher) Unmatched() []string {
unmatched := make([]string, 0, len(m.unmatchedReqs))
for k := range m.unmatchedReqs {
unmatched = append(unmatched, k)
}
return unmatched
}
// KubernetesResourceMatcher matches a role against a Kubernetes Resource.
// Kind is must be stricly equal but namespace and name allow wildcards.
type KubernetesResourceMatcher struct {
resource types.KubernetesResource
isClusterWideResource bool
}
// NewKubernetesResourceMatcher creates a KubernetesResourceMatcher that checks
// whether the role's KubeResources match the specified condition.
func NewKubernetesResourceMatcher(resource types.KubernetesResource, isClusterWideResource bool) *KubernetesResourceMatcher {
return &KubernetesResourceMatcher{
resource: resource,
isClusterWideResource: isClusterWideResource,
}
}
// Match matches a Kubernetes Resource against provided role and condition.
func (m *KubernetesResourceMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
result, err := utils.KubeResourceMatchesRegex(m.resource, m.isClusterWideResource, role.GetKubeResources(condition), condition)
return result, trace.Wrap(err)
}
// String returns the matcher's string representation.
func (m *KubernetesResourceMatcher) String() string {
return fmt.Sprintf("KubernetesResourceMatcher(Resource=%v)", m.resource)
}
// NewKubernetesClusterLabelMatcher creates a RoleMatcher that checks whether a role's
// Kubernetes service labels match.
func NewKubernetesClusterLabelMatcher(clustersLabels map[string]string, username string, userTraits wrappers.Traits) RoleMatcher {
return &kubernetesClusterLabelMatcher{clusterLabels: clustersLabels, username: username, userTraits: userTraits}
}
// Match matches a Kubernetes cluster labels against a role.
func (l *kubernetesClusterLabelMatcher) Match(role types.Role, typ types.RoleConditionType) (bool, error) {
labelMatchers, err := l.getKubeLabelMatchers(role, typ)
if err != nil {
return false, trace.Wrap(err)
}
ok, _, err := CheckLabelsMatch(typ, labelMatchers, l.username, l.userTraits, label.MapLabelGetter(l.clusterLabels), false)
return ok, trace.Wrap(err)
}
// getKubeLabelMatchers returns kubernetes_labels based on resource version and role type.
func (l kubernetesClusterLabelMatcher) getKubeLabelMatchers(role types.Role, typ types.RoleConditionType) (types.LabelMatchers, error) {
labelMatchers, err := role.GetLabelMatchers(typ, types.KindKubernetesCluster)
if err != nil {
return types.LabelMatchers{}, trace.Wrap(err)
}
// After the introduction of https://github.com/gravitational/teleport/pull/9759 the
// kubernetes_labels started to be respected. Former role behavior evaluated deny rules
// even if the kubernetes_labels was empty. To preserve this behavior after respecting kubernetes label the label
// logic needs to be aligned.
// Default wildcard rules should be added to deny.kubernetes_labels if
// deny.kubernetes_labels is empty to ensure that deny rule will be evaluated
// even if kubernetes_labels are empty.
if labelMatchers.Empty() && typ == types.Deny {
labelMatchers.Labels = types.Labels{types.Wildcard: []string{types.Wildcard}}
}
return labelMatchers, nil
}
// AccessCheckable is the subset of types.Resource required for the RBAC checks.
type AccessCheckable interface {
GetKind() string
GetSubKind() string
GetName() string
GetMetadata() types.Metadata
GetLabel(key string) (value string, ok bool)
GetAllLabels() map[string]string
}
var rbacLogger = logutils.NewPackageLogger(teleport.ComponentKey, teleport.ComponentRBAC)
// resourceRequiresLabelMatching decides if a resource requires label matching
// when making RBAC access decisions.
func resourceRequiresLabelMatching(r AccessCheckable) bool {
// Some resources do not need label matching when assessing whether the user
// should be granted access. Enable it by default, but turn it off in the
// special cases.
switch r.GetKind() {
case types.KindIdentityCenterAccount, types.KindIdentityCenterAccountAssignment:
return false
case types.KindApp, types.KindAppServer:
return r.GetSubKind() != types.KindIdentityCenterAccount
}
return true
}
// RoleGrantsResource reports whether role alone grants access to r, checking
// the namespace and label conditions only. It skips the MFA, device trust and
// lock checks, which apply to a whole role set rather than one role.
func RoleGrantsResource(role types.Role, r AccessCheckable, username string, traits wrappers.Traits) bool {
_, err := NewRoleSet(role).checkAccess(r, username, traits, AccessState{MFAVerified: true})
return err == nil
}
// checkAccess determines whether access should be granted to a resource based on the provided roles, resource
// attributes, user traits, access state (MFA, device trust, etc.), and optional matchers. If state.ReturnPreconditions
// is true, it returns a list of preconditions (e.g., MFA required) that must be satisfied for access. If
// state.ReturnPreconditions is false, it returns an error immediately if access is denied.
func (set RoleSet) checkAccess(
r AccessCheckable,
username string,
traits wrappers.Traits,
state AccessState,
matchers ...RoleMatcher,
) ([]*decisionpb.Precondition, error) {
// Note: logging in this function only happens in trace mode. This is because
// adding logging to this function (which is called on every resource returned
// by the backend) can slow down this function by 50x for large clusters!
ctx := context.Background()
logger := rbacLogger
isLoggingEnabled := logger.Handler().Enabled(ctx, logutils.TraceLevel)
if isLoggingEnabled {
logger = logger.With("resource_kind", r.GetKind(), "resource_name", r.GetName())
}
// Collect preconditions to return to the caller.
var preconds []*decisionpb.Precondition
// If the cluster requires per-session MFA and it hasn't been verified yet, add an MFA precondition or deny access early.
// If the legacy out-of-band MFA flow is allowed (see below) and MFA has already been verified for this session, skip this check.
//
// The legacy out-of-band MFA flow is allowed as long as TELEPORT_UNSTABLE_FORCE_IN_BAND_MFA is not set to "yes".
// When TELEPORT_UNSTABLE_FORCE_IN_BAND_MFA is set to "yes", only in-band MFA is allowed and enforced.
//
// TODO(cthach): Remove in v20.0 when the legacy out-of-band MFA flow is removed.
if state.MFARequired == MFARequiredAlways && (os.Getenv(teleport.EnvVarUnstableForceInBandMFA) == "yes" || !state.MFAVerified) {
// If the caller doesn't want preconditions returned, deny access early to avoid unnecessary work.
if !state.ReturnPreconditions {
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, cluster requires per-session MFA")
return nil, ErrSessionMFARequired
}
// Mark that MFA is required and continue evaluating access.
preconds = append(preconds, decisionpb.Precondition_builder{Kind: decisionpb.PreconditionKind_PRECONDITION_KIND_IN_BAND_MFA}.Build())
}
requiresLabelMatching := resourceRequiresLabelMatching(r)
namespace := types.ProcessNamespace(r.GetMetadata().Namespace)
// Additional message depending on kind of resource
// so there's more context on why the user might not have access.
additionalDeniedMessage := ""
switch r.GetKind() {
case types.KindDatabase:
additionalDeniedMessage = "Confirm database user and name."
case types.KindNode:
additionalDeniedMessage = "Confirm SSH login."
case types.KindKubernetesCluster:
additionalDeniedMessage = "Confirm Kubernetes user or group."
case types.KindWindowsDesktop:
additionalDeniedMessage = "Confirm Windows user."
case types.KindSAMLIdPServiceProvider:
additionalDeniedMessage = "Confirm app_labels."
}
// Check deny rules.
for _, role := range set {
matchNamespace, namespaceMessage := MatchNamespace(role.GetNamespaces(types.Deny), namespace)
if !matchNamespace {
continue
}
if requiresLabelMatching {
matchLabels, labelsMessage, err := checkRoleLabelsMatch(types.Deny, role, username, traits, r, isLoggingEnabled)
if err != nil {
return nil, trace.Wrap(err)
}
if matchLabels {
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, deny rule in role matched",
slog.String("role", role.GetName()),
slog.String("namespace_message", namespaceMessage),
slog.String("label_message", labelsMessage),
)
return nil, trace.AccessDenied("access to %v denied. User does not have permissions. %v",
r.GetKind(), additionalDeniedMessage)
}
} else {
logger.LogAttrs(ctx, logutils.TraceLevel, "Role label matching skipped")
}
// Deny rules are greedy on purpose. They will always match if
// at least one of the matchers returns true.
matchMatchers, matchersMessage, err := RoleMatchers(matchers).MatchAny(role, types.Deny)
if err != nil {
return nil, trace.Wrap(err)
}
if matchMatchers {
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, deny rule in role matched",
slog.String("role", role.GetName()),
slog.Any("matcher_message", matchersMessage),
)
return nil, trace.AccessDenied("access to %v denied. User does not have permissions. %v",
r.GetKind(), additionalDeniedMessage)
}
}
// MFA checks can be bypassed if either:
// 1. The cluster doesn't require per-session MFA (MFARequiredNever), OR
// 2. Legacy out-of-band MFA has already been verified for the session AND
// a. The legacy out-of-band MFA flow is allowed (TELEPORT_UNSTABLE_FORCE_IN_BAND_MFA is not set to "yes") OR
// b. The caller doesn't want preconditions returned (state.ReturnPreconditions is false)
//
// Listing resources sets state.MFAVerified to true and state.ReturnPreconditions to false to allow bypassing MFA
// checks for resources that require per-session MFA. This is because listing resources is a read-only operation and
// MFA is not required to list resources, even if MFA is required to access the resource. The actual enforcement
// will happen at connection time, so this is not a concern from a security perspective.
//
// TODO(cthach): Remove in v20.0 when the legacy out-of-band MFA flow is removed.
bypassMFAChecks := state.MFARequired == MFARequiredNever ||
(state.MFAVerified && (os.Getenv(teleport.EnvVarUnstableForceInBandMFA) != "yes" || !state.ReturnPreconditions))
// TODO(codingllama): Consider making EnableDeviceVerification opt-out instead
// of opt-in.
deviceTrusted := !state.EnableDeviceVerification || state.DeviceVerified
var errs []error
allowed := false
// Check allow rules.
for _, role := range set {
matchNamespace, namespaceMessage := MatchNamespace(role.GetNamespaces(types.Allow), namespace)
if !matchNamespace {
if isLoggingEnabled {
errs = append(errs, trace.AccessDenied("role=%v, match(namespace=%v)",
role.GetName(), namespaceMessage))
}
continue
}
if requiresLabelMatching {
matchLabels, labelsMessage, err := checkRoleLabelsMatch(types.Allow, role, username, traits, r, isLoggingEnabled)
if err != nil {
return nil, trace.Wrap(err)
}
if !matchLabels {
if isLoggingEnabled {
errs = append(errs, trace.AccessDenied("role=%v, match(%s)",
role.GetName(), labelsMessage))
}
continue
}
} else {
logger.LogAttrs(ctx, logutils.TraceLevel, "Role label matching skipped for resource")
}
// Allow rules are not greedy. They will match only if all of the
// matchers return true.
matchMatchers, err := RoleMatchers(matchers).MatchAll(role, types.Allow)
if err != nil {
return nil, trace.Wrap(err)
}
if !matchMatchers {
if isLoggingEnabled {
errs = append(errs, fmt.Errorf("role=%v, match(matchers=%v)",
role.GetName(), matchers))
}
continue
}
// If we've reached this point, namespace, labels, and matchers all match.
//
// The following checks remain:
// 1. MFA verification (aka require_session_mfa)
// 2. Device verification (aka device_trust_mode)
//
// The more restrictive setting applies, so either the caller passes all
// (and gets an early exit) or we need to check every applicable role to
// ensure the access is permitted.
if bypassMFAChecks && deviceTrusted {
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource granted, allow rule in role matched",
slog.String("role", role.GetName()),
)
return deduplicateAndSortPreconditions(preconds), nil
}
// Check if MFA is required at the role-level.
if !bypassMFAChecks && role.GetOptions().RequireMFAType.IsSessionMFARequired() {
// If the caller doesn't want preconditions returned, deny access early to avoid unnecessary work.
if !state.ReturnPreconditions {
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, role requires per-session MFA",
slog.String("role", role.GetName()),
)
return nil, ErrSessionMFARequired
}
// Mark that MFA is required and continue evaluating access.
preconds = append(preconds, decisionpb.Precondition_builder{Kind: decisionpb.PreconditionKind_PRECONDITION_KIND_IN_BAND_MFA}.Build())
}
// Device verification.
if err := dtauthz.VerifyTrustedDeviceMode(
role.GetOptions().DeviceTrustMode,
dtauthz.VerifyTrustedDeviceModeParams{
IsTrustedDevice: deviceTrusted,
IsBot: state.IsBot,
AllowEmptyMode: true, // Empty mode on roles is equivalent to "off".
},
); err != nil {
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, role requires a trusted device",
slog.String("role", role.GetName()),
)
return nil, trace.Wrap(err)
}
// Current role allows access, but keep looking for a more restrictive
// setting.
allowed = true
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource granted, allow rule in role matched",
slog.String("role", role.GetName()),
)
}
if allowed {
return deduplicateAndSortPreconditions(preconds), nil
}
logger.LogAttrs(ctx, logutils.TraceLevel, "Access to resource denied, no allow rule matched",
slog.Any("errors", errs),
)
return nil, trace.AccessDenied("access to %v denied. User does not have permissions. %v",
r.GetKind(), additionalDeniedMessage)
}
func deduplicateAndSortPreconditions(preconds []*decisionpb.Precondition) []*decisionpb.Precondition {
// Deduplicate preconditions by kind.
preconds = slices.CompactFunc(
preconds, func(a, b *decisionpb.Precondition) bool {
return a.GetKind() == b.GetKind()
},
)
// Sort by kind for deterministic ordering during enforcement.
slices.SortFunc(
preconds,
func(a, b *decisionpb.Precondition) int {
return cmp.Compare(a.GetKind(), b.GetKind())
},
)
return preconds
}
// checkRoleLabelsMatch checks if the [role] matches the labels of [resource]
// for [condition].
// It considers both the role labels (<kind>_labels) and label expression
// (<kind>_labels_expression).
//
// Returns a match boolean, a debug message, and any unexpected error.
//
// If [condition] is types.Deny, the match is greedy, if either one matches it's
// considered a match.
//
// If [condition] is types.Allow, the match is not greedy, if either doesn't
// match it's not considered a match.
//
// If neither is set, it's not a match in either case.
func checkRoleLabelsMatch(
condition types.RoleConditionType,
role types.Role,
username string,
userTraits wrappers.Traits,
resource AccessCheckable,
debug bool,
) (bool, string, error) {
labelMatchers, err := role.GetLabelMatchers(condition, resource.GetKind())
if err != nil {
return false, "", trace.Wrap(err)
}
return CheckLabelsMatch(condition, labelMatchers, username, userTraits, resource, debug)
}
// CheckLabelsMatch checks if the [labelMatchers] match the labels of [resource]
// for [condition].
// It considers both [labelMatchers.Labels] and [labelMatchers.Expression].
//
// Returns a match boolean, a debug message, and any unexpected error.
//
// If [condition] is types.Deny, the match is greedy, if either one matches it's
// considered a match.
//
// If [condition] is types.Allow, the match is not greedy, if either doesn't
// match it's not considered a match.
//
// If neither is set, it's not a match in either case.
func CheckLabelsMatch(
condition types.RoleConditionType,
labelMatchers types.LabelMatchers,
username string,
userTraits wrappers.Traits,
resource label.LabelGetter,
debug bool,
) (bool, string, error) {
if labelMatchers.Empty() {
return false, "no label matchers or label expression", nil
}
var message string
labelsUnsetOrMatch, expressionUnsetOrMatch := true, true
if len(labelMatchers.Labels) > 0 {
match, msg, err := MatchLabelGetter(labelMatchers.Labels, resource)
if err != nil {
return false, "", trace.Wrap(err)
}
if debug {
message += "label=" + msg
}
// Deny rules are greedy, if either matches, it's a match.
if condition == types.Deny && match {
return true, message, nil
}
labelsUnsetOrMatch = match
}
if len(labelMatchers.Expression) > 0 {
match, msg, err := matchLabelExpression(labelMatchers.Expression, resource, username, userTraits)
if err != nil {
return false, "", trace.Wrap(err)
}
if debug {
message = strings.Join([]string{message, "expression=" + msg}, ", ")
}
// Deny rules are greedy, if either matches, it's a match.
if condition == types.Deny {
return match, message, nil
}
expressionUnsetOrMatch = match
}
if condition == types.Deny {
// Either branch would have returned if it was a match.
return false, message, nil
}
// Allow rules are not greedy, both must match if they are set.
return labelsUnsetOrMatch && expressionUnsetOrMatch, message, nil
}
func matchLabelExpression(labelExpression string, resource label.LabelGetter, username string, userTraits wrappers.Traits) (bool, string, error) {
parsedExpr, err := label.ParseExpression(labelExpression)
if err != nil {
return false, "", trace.Wrap(err)
}
match, err := parsedExpr.Evaluate(label.ExpressionEnv{
ResourceLabelGetter: resource,
Username: username,
UserTraits: userTraits,
})
if err != nil {
return false, "", trace.Wrap(err, "evaluating label expression %q", labelExpression)
}
if match {
return true, "matched", nil
}
return false, "no match", nil
}
// CanForwardAgents returns true if role set allows forwarding agents.
func (set RoleSet) CanForwardAgents() bool {
for _, role := range set {
if role.GetOptions().ForwardAgent.Value() {
return true
}
}
return false
}
// SSHPortForwardMode returns the SSHPortForwardMode permitted by a RoleSet. Port forwarding is implicitly allowed, but explicit denies take
// precedence of explicit allows when using SSHPortForwarding. The legacy PortForwarding field prefers explicit allows for backwards
// compatibility reasons, but is only evaluated in the absence of an SSHPortForwarding config on the same role.
func (set RoleSet) SSHPortForwardMode() decisionpb.SSHPortForwardMode {
var denyRemote, denyLocal, legacyDeny bool
legacyCanDeny := true
for _, role := range set {
config := role.GetOptions().SSHPortForwarding
// only consider legacy allows when config isn't provided on the same role
if config == nil {
//nolint:staticcheck // this field is preserved for backwards compatibility, but shouldn't be used going forward
if legacy := role.GetOptions().PortForwarding; legacy != nil {
if legacy.Value {
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_ON
}
legacyDeny = true
}
continue
}
if config.Remote != nil && config.Remote.Enabled != nil {
if !config.Remote.Enabled.Value {
denyRemote = true
}
// an explicit legacy deny is only possible if no explicit SSHPortForwarding config has been provided
legacyCanDeny = false
}
if config.Local != nil && config.Local.Enabled != nil {
if !config.Local.Enabled.Value {
denyLocal = true
}
// an explicit legacy deny is only possible if no explicit SSHPortForwarding config has been provided
legacyCanDeny = false
}
}
// enforcing implicit allow and preferring allow over explicit deny
switch {
case denyRemote && denyLocal:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_OFF
case legacyDeny && legacyCanDeny:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_OFF
case denyRemote:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_LOCAL
case denyLocal:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_REMOTE
default:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_ON
}
}
// CanPortForward returns true if the RoleSet allows both local and remote port forwarding.
func (set RoleSet) CanPortForward() bool {
return set.SSHPortForwardMode() == decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_ON
}
// RecordDesktopSession returns true if the role set has enabled desktop
// session recording. Recording is considered enabled if at least one
// role in the set has enabled it.
func (set RoleSet) RecordDesktopSession() bool {
for _, role := range set {
var bo *types.BoolOption
if role.GetOptions().RecordSession != nil {
bo = role.GetOptions().RecordSession.Desktop
}
if types.BoolDefaultTrue(bo) {
return true
}
}
return false
}
// DesktopClipboard returns true if the role set has enabled shared
// clipboard for desktop sessions. Clipboard sharing is disabled if
// one or more of the roles in the set has disabled it.
func (set RoleSet) DesktopClipboard() bool {
for _, role := range set {
if !types.BoolDefaultTrue(role.GetOptions().DesktopClipboard) {
return false
}
}
return true
}
// DesktopDirectorySharing returns true if the role set has directory sharing
// enabled. This setting is disabled if one or more of the roles in the set has
// disabled it.
func (set RoleSet) DesktopDirectorySharing() bool {
for _, role := range set {
if !types.BoolDefaultTrue(role.GetOptions().DesktopDirectorySharing) {
return false
}
}
return true
}
// MaybeCanReviewRequests attempts to guess if this RoleSet belongs
// to a user who should be submitting access reviews. Because not all rolesets
// are derived from statically assigned roles, this may return false positives.
func (set RoleSet) MaybeCanReviewRequests() bool {
for _, role := range set {
if !role.GetAccessReviewConditions(types.Allow).IsZero() {
// at least one nonzero allow directive exists for
// review submission.
return true
}
}
return false
}
// PermitX11Forwarding returns true if this RoleSet allows X11 Forwarding.
func (set RoleSet) PermitX11Forwarding() bool {
for _, role := range set {
if role.GetOptions().PermitX11Forwarding.Value() {
return true
}
}
return false
}
// CanCopyFiles returns true if the role set has enabled remote file
// operations via SCP or SFTP. Remote file operations are disabled if
// one or more of the roles in the set has disabled it.
func (set RoleSet) CanCopyFiles() bool {
for _, role := range set {
if !types.BoolDefaultTrue(role.GetOptions().SSHFileCopy) {
return false
}
}
return true
}
// GetWebTerminalClipboardMode returns the Web UI terminal clipboard mode from the role set.
func (set RoleSet) GetWebTerminalClipboardMode() types.WebTerminalClipboardMode {
var mode types.WebTerminalClipboardMode
for _, r := range set {
switch r.GetOptions().WebTerminalClipboardMode {
// Return immediately if any role has explicitly set the clipboard mode to no-copy, as that should take precedence over any unrestricted's.
case types.WebTerminalClipboardMode_WEB_TERMINAL_CLIPBOARD_MODE_NO_COPY:
return types.WebTerminalClipboardMode_WEB_TERMINAL_CLIPBOARD_MODE_NO_COPY
case types.WebTerminalClipboardMode_WEB_TERMINAL_CLIPBOARD_MODE_UNRESTRICTED:
mode = types.WebTerminalClipboardMode_WEB_TERMINAL_CLIPBOARD_MODE_UNRESTRICTED
}
}
return mode
}
// CanJoinSessions returns true if at least one role in the role set
// allows the user to join active sessions.
func (set RoleSet) CanJoinSessions() bool {
return slices.ContainsFunc(set, func(r types.Role) bool {
return len(r.GetSessionJoinPolicies()) > 0
})
}
// CertificateFormat returns the most permissive certificate format in a
// RoleSet.
func (set RoleSet) CertificateFormat() string {
var formats []string
for _, role := range set {
// get the certificate format for each individual role. if a role does not
// have a certificate format (like implicit roles) skip over it
certificateFormat := role.GetOptions().CertificateFormat
if certificateFormat == "" {
continue
}
formats = append(formats, certificateFormat)
}
// if no formats were found, return standard
if len(formats) == 0 {
return constants.CertificateFormatStandard
}
// sort the slice so the most permissive is the first element
sort.Slice(formats, func(i, j int) bool {
return certificatePriority(formats[i]) < certificatePriority(formats[j])
})
return formats[0]
}
// EnhancedRecordingSet returns the set of enhanced session recording
// events to capture for thi role set.
func (set RoleSet) EnhancedRecordingSet() map[string]bool {
m := make(map[string]bool)
// Loop over all roles and create a set of all options.
for _, role := range set {
for _, opt := range role.GetOptions().BPF {
m[opt] = true
}
}
return m
}
// certificatePriority returns the priority of the certificate format. The
// most permissive has lowest value.
func certificatePriority(s string) int {
switch s {
case teleport.CertificateFormatOldSSH:
return 0
case constants.CertificateFormatStandard:
return 1
default:
return 2
}
}
// CheckAgentForward checks if the role can request to forward the SSH agent
// for this user.
func (set RoleSet) CheckAgentForward(login string) error {
// check if we have permission to login and forward agent. we don't check
// for deny rules because if you can't forward an agent if you can't login
// in the first place.
for _, role := range set {
for _, l := range role.GetLogins(types.Allow) {
if role.GetOptions().ForwardAgent.Value() && l == login {
return nil
}
}
}
return trace.AccessDenied("%v can not forward agent for %v", set, login)
}
func (set RoleSet) String() string {
if len(set) == 0 {
return "user without assigned roles"
}
roleNames := make([]string, len(set))
for i, role := range set {
roleNames[i] = role.GetName()
}
return fmt.Sprintf("roles %v", strings.Join(roleNames, ","))
}
// GuessIfAccessIsPossible guesses if access is possible for an entire category
// of resources.
// It responds the question: "is it possible that there is a resource of this
// kind that the current user can access?".
// GuessIfAccessIsPossible is used, mainly, for UI decisions ("should the tab
// for resource X appear"?). Most callers should use CheckAccessToRule instead.
func (set RoleSet) GuessIfAccessIsPossible(ctx RuleContext, namespace string, resource string, verb string) error {
// "Where" clause are handled differently by the method:
// - "allow" rules have their "where" clause always match, as it's assumed
// that there could be a resource that matches it.
// - "deny" rules have their "where" clause always fail, as it's assumed that
// there could be a resource that passes it.
return set.checkAccessToRuleImpl(checkAccessParams{
ctx: ctx,
namespace: namespace,
resource: resource,
verb: verb,
allowWhere: boolParser(true), // always matches
denyWhere: boolParser(false), // never matches
})
}
type boolParser bool
func (p boolParser) Parse(string) (any, error) {
return predicate.BoolPredicate(func() bool {
return bool(p)
}), nil
}
// CheckAccessToRule checks if the RoleSet provides access in the given
// namespace to the specified resource and verb.
// silent controls whether the access violations are logged.
func (set RoleSet) CheckAccessToRule(ctx RuleContext, namespace string, resource string, verb string) error {
whereParser, err := NewWhereParser(
ctx,
// register can_view function if the resource is a session.
ConditionalOption(resource == types.KindSession, WithCanViewFunction()),
)
if err != nil {
return trace.Wrap(err)
}
return set.checkAccessToRuleImpl(checkAccessParams{
ctx: ctx,
namespace: namespace,
resource: resource,
verb: verb,
allowWhere: whereParser,
denyWhere: whereParser,
})
}
// GetKubeResources returns allowed and denied list of Kubernetes Resources configured in the RoleSet.
func (set RoleSet) GetKubeResources(cluster types.KubeCluster, username string, userTraits wrappers.Traits) (allowed, denied []types.KubernetesResource) {
for _, role := range set {
matchLabels, _, err := checkRoleLabelsMatch(types.Allow, role, username, userTraits, cluster, false)
if err != nil || !matchLabels {
continue
}
allowed = append(allowed, role.GetKubeResources(types.Allow)...)
}
for _, role := range set {
// deny rules are not checked for labels because they are greedy. It means that
// if there is a deny rule for a cluster, it will deny access to all resources
// in that cluster, regardless of kubernetes_resources (i.e. making them irrelevant).
// If the goal is to deny access to a specific resource, it should be done by collecting
// all kube resources in deny rules and ignoring if the role matches or not
// the cluster (i.e. no labels check).
denied = append(denied, role.GetKubeResources(types.Deny)...)
}
return deduplicateKubeResources(allowed), deduplicateKubeResources(denied)
}
func deduplicateKubeResources(resources []types.KubernetesResource) []types.KubernetesResource {
allKeys := setutils.New[string]()
copy := make([]types.KubernetesResource, 0, len(resources))
for _, item := range resources {
key := item.String()
if !allKeys.Contains(key) {
allKeys.Add(key)
copy = append(copy, item)
}
}
return copy
}
type checkAccessParams struct {
ctx RuleContext
namespace string
resource string
verb string
allowWhere, denyWhere predicate.Parser
}
type accessExplicitlyDenied struct {
inner error
}
// AccessExplicitlyDenied is an error type that indicates an AccessDenied error
// where a deny rule matched and access is explicitly denied, in contrast to
// cases where there is no matching deny or allow rule and access is only
// implicitly denied.
func AccessExplicitlyDenied(inner error) error {
return &accessExplicitlyDenied{inner}
}
// IsAccessExplicitlyDenied returns true if any of the errors in err's chain is
// an AccessExplicitlyDenied error.
func IsAccessExplicitlyDenied(err error) bool {
var target *accessExplicitlyDenied
return errors.As(err, &target)
}
func (a *accessExplicitlyDenied) Error() string {
return a.inner.Error()
}
func (a *accessExplicitlyDenied) Unwrap() error {
return a.inner
}
func (set RoleSet) checkAccessToRuleImpl(p checkAccessParams) (err error) {
ctx := context.Background()
// Every unknown error, which could be due to a bad role or an expression
// that can't parse, should be considered an explicit denial.
explicitDeny := true
defer func() {
if explicitDeny && err != nil {
err = AccessExplicitlyDenied(err)
}
}()
actionsParser, err := NewActionsParser(p.ctx)
if err != nil {
return trace.Wrap(err)
}
// check deny: a single match on a deny rule prohibits access
for _, role := range set {
matchNamespace, _ := MatchNamespace(role.GetNamespaces(types.Deny), types.ProcessNamespace(p.namespace))
if matchNamespace {
matched, err := MakeRuleSet(role.GetRules(types.Deny)).Match(p.denyWhere, actionsParser, p.resource, p.verb)
if err != nil {
return trace.Wrap(err)
}
if matched {
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access denied, deny rule matched",
slog.String("verb", p.verb),
slog.String("resource", p.resource),
slog.String("namespace", p.namespace),
slog.String("role", role.GetName()),
)
return trace.AccessDenied("access denied to perform action %q on %q", p.verb, p.resource)
}
}
}
// check allow: if rule matches, grant access to resource
for _, role := range set {
matchNamespace, _ := MatchNamespace(role.GetNamespaces(types.Allow), types.ProcessNamespace(p.namespace))
if matchNamespace {
match, err := MakeRuleSet(role.GetRules(types.Allow)).Match(p.allowWhere, actionsParser, p.resource, p.verb)
if err != nil {
return trace.Wrap(err)
}
if match {
return nil
}
}
}
rbacLogger.LogAttrs(ctx, logutils.TraceLevel, "Access denied, no allow rule matched",
slog.String("verb", p.verb),
slog.String("resource", p.resource),
slog.String("namespace", p.namespace),
slog.Any("set", set),
)
// At this point no deny rule has matched and there are no more unknown
// errors, so this is only an implicit denial.
explicitDeny = false
return trace.AccessDenied("access denied to perform action %q on %q", p.verb, p.resource)
}
// ExtractConditionForIdentifier returns a restrictive filter expression
// for list queries based on the rules' `where` conditions.
func (set RoleSet) ExtractConditionForIdentifier(ctx RuleContext, namespace, resource, verb, identifier string) (*types.WhereExpr, error) {
parser, err := newParserForIdentifierSubcondition(ctx, identifier)
if err != nil {
return nil, trace.Wrap(err)
}
parseWhere := func(rule types.Rule) (types.WhereExpr, error) {
if rule.Where == "" {
return types.WhereExpr{Literal: true}, nil
}
out, err := parser.Parse(rule.Where)
if err != nil {
return types.WhereExpr{}, trace.Wrap(err)
}
expr, ok := out.(types.WhereExpr)
if !ok {
return types.WhereExpr{}, trace.BadParameter("invalid type %T when extracting identifier subcondition from %q", out, rule.Where)
}
return expr, nil
}
// Gather identifier-related subconditions from the deny rules
// and concatenate their negations by AND.
var denyCond *types.WhereExpr
for _, role := range set {
matchNamespace, _ := MatchNamespace(role.GetNamespaces(types.Deny), types.ProcessNamespace(namespace))
if !matchNamespace {
continue
}
rules := MakeRuleSet(role.GetRules(types.Deny))
for _, rule := range rules[resource] {
if !rule.HasVerb(verb) && !rule.HasVerb(types.Wildcard) {
continue
}
expr, err := parseWhere(rule)
if err != nil {
return nil, trace.Wrap(err)
}
if b, ok := expr.Literal.(bool); ok {
if b {
return nil, trace.AccessDenied("access denied to perform action %q on %q", verb, resource)
}
continue
}
negated := types.WhereExpr{Not: &expr}
if denyCond == nil {
denyCond = &negated
} else {
denyCond = &types.WhereExpr{And: types.WhereExpr2{L: denyCond, R: &negated}}
}
}
}
// Gather identifier-related subconditions from the allow rules
// and concatenate by OR.
var allowCond *types.WhereExpr
for _, role := range set {
matchNamespace, _ := MatchNamespace(role.GetNamespaces(types.Allow), types.ProcessNamespace(namespace))
if !matchNamespace {
continue
}
rules := MakeRuleSet(role.GetRules(types.Allow))
for _, rule := range rules[resource] {
if !rule.HasVerb(verb) && !rule.HasVerb(types.Wildcard) {
continue
}
expr, err := parseWhere(rule)
if err != nil {
return nil, trace.Wrap(err)
}
if b, ok := expr.Literal.(bool); ok {
if b {
return denyCond, nil
}
continue
}
if allowCond == nil {
allowCond = &expr
} else {
allowCond = &types.WhereExpr{Or: types.WhereExpr2{L: allowCond, R: &expr}}
}
}
}
if denyCond == nil {
if allowCond == nil {
return nil, trace.AccessDenied("access denied to perform action %q on %q", verb, resource)
}
return allowCond, nil
}
return &types.WhereExpr{And: types.WhereExpr2{L: denyCond, R: allowCond}}, nil
}
// SearchAsRolesOption is a functional option for filtering SearchAsRoles.
type SearchAsRolesOption func(role types.Role) bool
// GetSearchAsRoles returns all SearchAsRoles for this RoleSet.
func (set RoleSet) GetAllowedSearchAsRoles(allowFilters ...SearchAsRolesOption) []string {
denied := make(map[string]struct{})
var allowed []string
for _, role := range set {
for _, d := range role.GetSearchAsRoles(types.Deny) {
denied[d] = struct{}{}
}
}
for _, role := range set {
if slices.ContainsFunc(allowFilters, func(filter SearchAsRolesOption) bool {
return !filter(role)
}) {
// Don't consider this base role if it's filtered out.
continue
}
for _, a := range role.GetSearchAsRoles(types.Allow) {
if _, isDenied := denied[a]; isDenied {
continue
}
allowed = append(allowed, a)
}
}
return apiutils.Deduplicate(allowed)
}
type gk struct{ group, kind string }
// noramlize the give kube kind. Maps legacy values to plural+group, trim the kube: prefix.
// Returns <kind>[.<group>].
func normalizeKubernetesKind(in string) (out gk) {
// Check if we have a legacy kind.
out.group = types.KubernetesResourcesV7KindGroups[in]
out.kind = types.KubernetesResourcesKindsPlurals[in]
if out.kind == "" {
switch {
case in == types.KindKubeNamespace:
out.kind = "namespaces"
return out
case strings.HasPrefix(in, types.AccessRequestPrefixKindKubeNamespaced):
out.kind = strings.TrimPrefix(in, types.AccessRequestPrefixKindKubeNamespaced)
case strings.HasPrefix(in, types.AccessRequestPrefixKindKubeClusterWide):
out.kind = strings.TrimPrefix(in, types.AccessRequestPrefixKindKubeClusterWide)
// Subset if the two first used in search. Must be last.
case strings.HasPrefix(in, types.AccessRequestPrefixKindKube):
out.kind = strings.TrimPrefix(in, types.AccessRequestPrefixKindKube)
}
}
if out.group != "" { // If we have a group, we are dealing with legacy value, we have the noramlized version.
return out
}
// Otherwise, parse the group from the trimmed input.
if i := strings.Index(out.kind, "."); i != -1 {
out.group = out.kind[i+1:]
out.kind = out.kind[:i]
return out
}
return out
}
// matchRequestKubernetesResources checks if the input matches the reference
// based on the condition type.
//
// Similar logic as utils.KubeResourceMatchesRegex(), but with support for wildcard input
// and without support for verbs/names/namespaces.
//
// Examples:
// Request: *.apps Deny: deployments.apps -> match.
// Request: *.apps Deny: deployments.* -> match. (*.apps could be deployments.apps which matches deployments.*)
// Request: *.apps Deny: *.* -> match.
// Request: deployments.* Deny: deployments.apps -> match.
// Request: deployments.* Deny: deployments.* -> match.
// Request: deployments.* Deny: *.* -> match.
// Request: *.* Deny: deployments.apps -> match.
// Request: *.* Deny: deployments.* -> match.
// Request: *.* Deny: *.* -> match.
func matchRequestKubernetesResources(input gk, reference types.RequestKubernetesResource, cond types.RoleConditionType) bool {
// If we have an exact match, we are done.
if input.kind == reference.Kind && input.group == reference.APIGroup {
return true
}
// If the reference is a wildcard and the input kube_cluster, we don't match allow, but we match deny.
// Ref:
// https://github.com/gravitational/teleport/blob/master/rfd/0183-access-request-kube-resource-allow-list.md#as-an-admin-i-want-to-require-users-to-request-for-kubernetes-subresources-instead-of-the-whole-kubernetes-cluster
if reference.Kind == types.Wildcard && input.kind == types.KindKubernetesCluster {
return cond == types.Deny
}
if cond == types.Allow {
// In allow mode, if the reference kind is not a wildcard and doesn't match exactly, we reject.
if reference.Kind != types.Wildcard && input.kind != reference.Kind {
return false
}
// If the reference api group is a wildcard or is an exact match, we are done.
if reference.APIGroup == types.Wildcard || input.group == reference.APIGroup {
return true
}
// Otherwise, attempt to match the api group pattern.
ok, _ := utils.MatchString(input.group, reference.APIGroup)
return ok
}
// In deny mode, we reject only if both input/ref are not wildcard and are not equal.
if reference.Kind != types.Wildcard && input.kind != types.Wildcard && input.kind != reference.Kind {
return false
}
// If there is no conflict on the kind, check the group. As we support pattern matching, check both sides.
ok1, _ := utils.MatchString(input.group, reference.APIGroup)
ok2, _ := utils.MatchString(reference.APIGroup, input.group)
return ok1 || ok2
}
// GetAllowedSearchAsRolesForKubeResourceKind returns all of the allowed SearchAsRoles
// that allowed requesting to the requested Kubernetes resource kind.
func (set RoleSet) GetAllowedSearchAsRolesForKubeResourceKind(requestedKubeResourceKind string) []string {
// Return no results if encountering any denies since its globally matched.
for _, role := range set {
for _, kr := range role.GetRequestKubernetesResources(types.Deny) {
if matchRequestKubernetesResources(normalizeKubernetesKind(requestedKubeResourceKind), kr, types.Deny) {
return nil
}
}
}
return set.GetAllowedSearchAsRoles(WithAllowedKubernetesResourceKindFilter(requestedKubeResourceKind))
}
// WithAllowedKubernetesResourceKindFilter returns a SearchAsRolesOption func
// that will check that the requestedKubeResourceKind exists in the allow list
// for the current role.
func WithAllowedKubernetesResourceKindFilter(requestedKubeResourceKind string) SearchAsRolesOption {
return func(role types.Role) bool {
allowed := role.GetAccessRequestConditions(types.Allow).KubernetesResources
// Any kind is allowed if nothing was configured.
if len(allowed) == 0 {
return true
}
for _, kr := range role.GetRequestKubernetesResources(types.Allow) {
if matchRequestKubernetesResources(normalizeKubernetesKind(requestedKubeResourceKind), kr, types.Allow) {
return true
}
}
return false
}
}
// GetAllowedPreviewAsRoles returns all PreviewAsRoles for this RoleSet.
func (set RoleSet) GetAllowedPreviewAsRoles() []string {
denied := make(map[string]struct{})
var allowed []string
for _, role := range set {
for _, d := range role.GetPreviewAsRoles(types.Deny) {
denied[d] = struct{}{}
}
}
for _, role := range set {
for _, a := range role.GetPreviewAsRoles(types.Allow) {
if _, ok := denied[a]; !ok {
allowed = append(allowed, a)
}
}
}
return apiutils.Deduplicate(allowed)
}
// GetCreateDatabaseUserMode returns the create database user mode of the rule
// set.
func (set RoleSet) GetCreateDatabaseUserMode() types.CreateDatabaseUserMode {
var mode types.CreateDatabaseUserMode
for _, r := range set {
if roleMode := r.GetCreateDatabaseUserMode(); roleMode > mode {
mode = roleMode
}
}
return mode
}
// AccessState holds state for the present access attempt, including both
// cluster settings and user state (MFA, device trust, etc).
type AccessState struct {
// MFARequired determines whether a user's MFA requirement dynamically changes
// based on their active role (per-role), or is static across all roles
// (always/never).
MFARequired MFARequired
// MFAVerified is set when MFA has been verified by the caller.
MFAVerified bool
// EnableDeviceVerification enables device verification in access checks.
// It's recommended to set this in tandem with DeviceVerified, so device
// checks are easier to reason about and have a proper chance of succeeding.
// Used for role-based device mode checks.
// Defaults to false for backwards compatibility.
EnableDeviceVerification bool
// DeviceVerified is true if the user certificate contains all required
// device extensions.
// A value of true enables the caller to clear device trust checks.
// It's recommended to set this in tandem with EnableDeviceVerification.
// See [dtauthz.IsTLSDeviceVerified] and [dtauthz.IsSSHDeviceVerified].
DeviceVerified bool
// IsBot determines whether the user certificate belongs to a bot. It's used
// when deciding whether to enforce device verification.
IsBot bool
// ReturnPreconditions, when set to true, causes access checks to return a set of preconditions (such as MFA or
// device verification requirements) instead of immediately returning an access error. This allows callers to
// programmatically determine what additional steps are required for access, rather than failing outright.
ReturnPreconditions bool
}
// MFARequired determines when MFA is required for a user to access a resource.
type MFARequired string
const (
// MFARequiredNever means that MFA is never required for any sessions started by this user.
// This means that it is not required by the cluster auth preference or any of the user's roles.
MFARequiredNever MFARequired = "never"
// MFARequiredAlways means that MFA is required for all sessions started by a user. This either
// means that the cluster auth preference requires per-session MFA, or all of the user's roles require
// per-session MFA
MFARequiredAlways MFARequired = "always"
// MFARequiredPerRole means that MFA requirement is based on which of the user's roles
// provides access to the session in question.
MFARequiredPerRole MFARequired = "per-role"
)
// UserSessionRoleNotFoundErrorMsg is added to "role not found" errors when they occur
// during user session roles validation. This allows the Web UI to distinguish between
// a user session role lookup error (which should prompt the user to re-login) vs. other role lookup
// failures.
// Keep in sync with teleport/src/services/api/api.ts(isUserSessionRoleNotFoundError)
const UserSessionRoleNotFoundErrorMsg = "user session role not found"
// knownAppResourceFields lists the JSON field names this version understands
// on an app_resources rule, derived from the AppResource message so a new
// proto field extends it automatically.
var knownAppResourceFields = func() map[string]struct{} {
fields := make(map[string]struct{})
for f := range reflect.TypeFor[types.AppResource]().Fields() {
name, _, _ := strings.Cut(f.Tag.Get("json"), ",")
if name != "" && name != "-" {
fields[name] = struct{}{}
}
}
return fields
}()
// denyAppAccessForUnknownFields empties any v9 allow app_resources rule whose
// stored JSON carried a field this version does not recognize, so the rule
// grants no access. Worst case is over-deny.
func denyAppAccessForUnknownFields(role *types.RoleV6, raw []byte) {
if role.Version != types.V9 {
return
}
for i := range role.Spec.Allow.AppResources {
for _, key := range jsoniter.Get(raw, "spec", "allow", "app_resources", i).Keys() {
if _, known := knownAppResourceFields[key]; !known {
role.Spec.Allow.AppResources[i] = types.AppResource{}
break
}
}
}
}
// UnmarshalRole unmarshals the Role resource from JSON.
func UnmarshalRole(bytes []byte, opts ...MarshalOption) (types.Role, error) {
return UnmarshalRoleV6(bytes, opts...)
}
// UnmarshalRoleV6 unmarshals the RoleV6 resource from JSON.
func UnmarshalRoleV6(bytes []byte, opts ...MarshalOption) (*types.RoleV6, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
version := jsoniter.Get(bytes, "version").ToString()
switch version {
// these are all backed by the same shape of data, they just have different semantics and defaults
case types.V3, types.V4, types.V5, types.V6, types.V7, types.V8, types.V9:
default:
return nil, trace.BadParameter("role version %q is not supported", version)
}
var role types.RoleV6
if err := utils.FastUnmarshal(bytes, &role); err != nil {
return nil, trace.BadParameter("%s", err)
}
if role.Version != version {
return nil, trace.BadParameter("inconsistent version in role data, got %q and %q", role.Version, version)
}
if cfg.DisallowUnknown {
if err := checkUnknownFields(bytes); err != nil {
return nil, trace.Wrap(err)
}
}
if err := CheckAndSetDefaults(&role); err != nil {
return nil, trace.Wrap(err)
}
denyAppAccessForUnknownFields(&role, bytes)
if cfg.Revision != "" {
role.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
role.SetExpiry(cfg.Expires)
}
return &role, nil
}
// MarshalRole marshals the Role resource to JSON.
func MarshalRole(role types.Role, opts ...MarshalOption) ([]byte, error) {
if err := CheckAndSetDefaults(role); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch role := role.(type) {
case *types.RoleV6:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, role))
default:
return nil, trace.BadParameter("unrecognized role version %T", role)
}
}
// checkUnknownFields rejects JSON with fields not defined in the RoleV6 struct.
func checkUnknownFields(data []byte) error {
var unused types.RoleV6
dec := json.NewDecoder(bytes.NewReader(data))
dec.DisallowUnknownFields()
if err := dec.Decode(&unused); err != nil {
return trace.BadParameter("role has unknown or misspelled fields: %v", err)
}
return nil
}
// AuthPreferenceGetter defines an interface for getting the authentication
// preferences.
type AuthPreferenceGetter interface {
// GetAuthPreference fetches the cluster authentication preferences.
GetAuthPreference(ctx context.Context) (types.AuthPreference, error)
}
// AccessStateFromSSHIdentity populates access state based on user's SSH
// identity and auth preference.
func AccessStateFromSSHIdentity(ctx context.Context, ident *sshca.Identity, checker AccessChecker, authPrefGetter AuthPreferenceGetter) (AccessState, error) {
authPref, err := authPrefGetter.GetAuthPreference(ctx)
if err != nil {
return AccessState{}, trace.Wrap(err)
}
state := checker.GetAccessState(authPref)
state.MFAVerified = ident.MFAVerified != ""
// Certain hardware-key based private key policies are treated as MFA verification.
if ident.PrivateKeyPolicy.MFAVerified() {
state.MFAVerified = true
}
state.EnableDeviceVerification = true
state.DeviceVerified = dtauthz.IsSSHDeviceVerified(ident)
state.IsBot = ident.IsBot()
return state, nil
}
// AccessStateFromTLSIdentity populates access state based on user's TLS
// identity and auth preference.
func AccessStateFromTLSIdentity(ctx context.Context, ident *tlsca.Identity, checker AccessChecker, authPrefGetter AuthPreferenceGetter) (AccessState, error) {
authPref, err := authPrefGetter.GetAuthPreference(ctx)
if err != nil {
return AccessState{}, trace.Wrap(err)
}
state := checker.GetAccessState(authPref)
state.MFAVerified = ident.MFAVerified != ""
// Certain hardware-key based private key policies are treated as MFA verification.
if ident.PrivateKeyPolicy.MFAVerified() {
state.MFAVerified = true
}
state.EnableDeviceVerification = true
state.DeviceVerified = dtauthz.IsTLSDeviceVerified(&ident.DeviceExtensions)
state.IsBot = ident.IsBot()
return state, nil
}
// MCPToolMatcher matches a role against MCP tool.
type MCPToolMatcher struct {
Name string
}
// Match matches MCP tool name against provided role and condition.
func (m *MCPToolMatcher) Match(role types.Role, condition types.RoleConditionType) (bool, error) {
mcpSpec := role.GetMCPPermissions(condition)
if mcpSpec == nil {
return false, nil
}
match, err := utils.SliceMatchesRegex(m.Name, mcpSpec.Tools)
return match, trace.Wrap(err)
}
// String returns the matcher's string representation.
func (m *MCPToolMatcher) String() string {
return fmt.Sprintf("MCPToolMatcher(Name=%v)", m.Name)
}
/*
* 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"
"crypto/x509"
"crypto/x509/pkix"
"encoding/base64"
"encoding/xml"
"errors"
"log/slog"
"net/http"
"os"
"slices"
"strings"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
saml2 "github.com/russellhaering/gosaml2"
samltypes "github.com/russellhaering/gosaml2/types"
dsig "github.com/russellhaering/goxmldsig"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
)
type SAMLConnectorGetter interface {
GetSAMLConnector(ctx context.Context, id string, withSecrets bool) (types.SAMLConnector, error)
GetSAMLConnectorWithValidationOptions(ctx context.Context, id string, withSecrets bool, opts ...types.SAMLConnectorValidationOption) (types.SAMLConnector, error)
}
type samlConnectorGetter func() (types.SAMLConnector, error)
const (
// ErrMsgHowToFixMissingPrivateKey is the error message displayed when validation of the signing key pair for secret refill fails.
ErrMsgHowToFixMissingPrivateKey = "You must either specify the signing key pair (obtain the existing one with `tctl get saml --with-secrets`) or let Teleport generate a new one (remove signing_key_pair in the resource you're trying to create)."
// ErrMsgHowToFixMissingOAuthCreds is the error message displayed when validation of the OAuth credentials for secret refill fails.
ErrMsgHowToFixMissingOAuthCreds = "You must specify the OAuth credentials (obtain the existing one with `tctl get saml --with-secrets`)."
)
// ErrFailedToFetchOrParseEntityDescriptor is returned on any error while downloading and parsing
// SAML entity descriptor if entity_descriptor_url is set during the connector validation.
var ErrFailedToFetchOrParseEntityDescriptor = &trace.BadParameterError{Message: "failed to fetch or parse entity descriptor"}
// ValidateSAMLConnector validates the SAMLConnector and sets default values.
// If a remote to fetch roles is specified, roles will be validated to exist.
func ValidateSAMLConnector(sc types.SAMLConnector, rg RoleGetter, opts ...types.SAMLConnectorValidationOption) error {
ctx := context.TODO()
options := types.NewSAMLConnectorValidationOptions(opts)
log := slog.With(teleport.ComponentKey, teleport.ComponentSAML, "saml_connector", sc.GetName())
// WARNING: this validation runs on the read and write paths, which means it
// is NOT safe to add new validations here that could invalidate existing
// connectors.
if err := CheckAndSetDefaults(sc); err != nil {
return trace.Wrap(err)
}
if creds := sc.GetCredentials(); creds != nil {
if err := creds.Validate(options.WithSecrets); err != nil {
return trace.Wrap(err)
}
}
if skp := sc.GetSigningKeyPair(); skp != nil {
if err := skp.Validate(options.WithSecrets); err != nil {
return trace.Wrap(err)
}
}
rawEntityDescriptor, entityDescriptor, err := getEntityDescriptor(ctx, getEntityDescriptorParams{
Log: log,
MFA: false,
Connector: sc,
Options: options,
})
if err != nil {
return trace.Wrap(err)
}
sc.SetEntityDescriptor(rawEntityDescriptor)
if ed := entityDescriptor; ed != nil {
sc.SetIssuer(ed.EntityID)
if ed.IDPSSODescriptor != nil && len(ed.IDPSSODescriptor.SingleSignOnServices) > 0 {
metadataSsoUrl := ed.IDPSSODescriptor.SingleSignOnServices[0].Location
if sc.GetSSO() != "" && sc.GetSSO() != metadataSsoUrl {
log.WarnContext(ctx,
"Connector has set SSO URL, but it does not match the one found in IDP metadata. Overwriting with the IDP metadata SSO URL.",
"connector_sso_url", sc.GetSSO(), "idp_metadata_sso_url", metadataSsoUrl,
)
}
sc.SetSSO(metadataSsoUrl)
}
}
if sc.GetIssuer() == "" {
return trace.BadParameter("no issuer or entityID set, either set issuer as a parameter or via entity_descriptor spec")
}
if sc.GetSSO() == "" {
return trace.BadParameter("no SSO set either explicitly or via entity_descriptor spec")
}
if err := validateAssertionConsumerServicesEndpoint(sc.GetSSO()); err != nil {
return trace.Wrap(err)
}
if sc.GetSigningKeyPair() == nil {
keyPEM, certPEM, err := utils.GenerateSelfSignedSigningCert(pkix.Name{
Organization: []string{"Teleport OSS"},
CommonName: "teleport.localhost.localdomain",
}, nil, 10*365*24*time.Hour)
if err != nil {
return trace.Wrap(err)
}
sc.SetSigningKeyPair(&types.AsymmetricKeyPair{
PrivateKey: string(keyPEM),
Cert: string(certPEM),
})
}
if options.WithAttributesToRoles {
if len(sc.GetAttributesToRoles()) == 0 {
return trace.BadParameter("attributes_to_roles is empty, authorization with connector would never assign any roles")
}
}
if rg != nil {
for _, mapping := range sc.GetAttributesToRoles() {
for _, role := range mapping.Roles {
if utils.ContainsExpansion(role) {
// Role is a template so we cannot check for existence of that literal name.
continue
}
_, err := rg.GetRole(ctx, role)
switch {
case trace.IsNotFound(err):
return trace.BadParameter("role %q specified in attributes_to_roles not found", role)
case err != nil:
return trace.Wrap(err)
}
}
}
}
preferredRequestBinding := sc.GetPreferredRequestBinding()
if preferredRequestBinding != "" {
if !slices.Contains(types.SAMLRequestBindingValues, preferredRequestBinding) {
return trace.BadParameter("invalid preferred_request_binding value. It can be one of %q", types.SAMLRequestBindingValues)
}
}
// Validate MFA settings.
if mfa := sc.GetMFASettings(); mfa != nil {
var mfaEntityDescriptor *samltypes.EntityDescriptor
if mfa.EntityDescriptorUrl != "" && mfa.EntityDescriptorUrl == sc.GetEntityDescriptorURL() {
// we got the entity descriptor above, skip the redundant round trip.
mfa.EntityDescriptor = rawEntityDescriptor
mfaEntityDescriptor = entityDescriptor
}
if mfaEntityDescriptor == nil {
var rawMFAEntityDescriptor string
rawMFAEntityDescriptor, mfaEntityDescriptor, err = getEntityDescriptor(ctx, getEntityDescriptorParams{
Log: log,
MFA: true,
Connector: sc,
Options: options,
})
if err != nil {
return trace.Wrap(err)
}
mfa.EntityDescriptor = rawMFAEntityDescriptor
}
if ed := mfaEntityDescriptor; ed != nil {
mfa.Issuer = ed.EntityID
if ed.IDPSSODescriptor != nil && len(ed.IDPSSODescriptor.SingleSignOnServices) > 0 {
mfa.Sso = ed.IDPSSODescriptor.SingleSignOnServices[0].Location
}
}
if preferredRequestBinding == types.SAMLRequestHTTPPostBinding {
log.WarnContext(ctx,
"SSO MFA does not support http-post binding request and will use the default http-redirect binding request",
"preferred_request_binding", preferredRequestBinding,
)
}
sc.SetMFASettings(mfa)
}
// TODO(nixpig): Add validation for when EntraIDGroupsProvider is present and not disbled then
// a valid authentication mechanism must be available. Proposed solution is to inject the
// token provider and ensure a valid azcore.TokenCredential can be resolved.
log.DebugContext(ctx, "Connector validated",
"sso", sc.GetSSO(),
"issuer", sc.GetIssuer(),
"acs", sc.GetAssertionConsumerService(),
)
return nil
}
type getEntityDescriptorParams struct {
Log *slog.Logger
MFA bool
Connector types.SAMLConnector
Options types.SAMLConnectorValidationOptions
}
func getEntityDescriptor(ctx context.Context, params getEntityDescriptorParams) (raw string, _ *samltypes.EntityDescriptor, err error) {
connector := params.Connector
rawEntityDescriptor, url := connector.GetEntityDescriptor(), connector.GetEntityDescriptorURL()
if params.MFA {
mfa := connector.GetMFASettings()
if mfa == nil {
return "", nil, trace.BadParameter("MFA set and MFA settings missing in the connector (this is a bug)")
}
rawEntityDescriptor, url = mfa.EntityDescriptor, mfa.EntityDescriptorUrl
}
log := params.Log
switch {
case params.MFA && url != "":
log = log.With("mfa_entity_descriptor_url", url)
case url != "":
log = log.With("entity_descriptor_url", url)
}
// Sanitize the error message to mitigate potential SSRF attacks attempted by SAML admins.
defer func() {
if url == "" || params.Options.NoFollowURLs || err == nil {
return
}
if params.MFA {
log.ErrorContext(ctx, "Failed to fetch or parse SAML MFA entity descriptor", "error", err)
} else {
log.ErrorContext(ctx, "Failed to fetch or parse SAML entity descriptor", "error", err)
}
err = trace.Wrap(ErrFailedToFetchOrParseEntityDescriptor)
}()
if url != "" && !params.Options.NoFollowURLs {
var checkRedirect func(req *http.Request, via []*http.Request) error
// TODO(kopiczko): Remove this env var after Jul 2027 (one year since introduced) if no issue is reported.
if disableCheckRedirect, _ := apiutils.ParseBool(os.Getenv(teleport.EnvVarUnstableDisableSAMLRedirectDowngradeCheck)); disableCheckRedirect {
log.DebugContext(ctx, "Redirect HTTPS downgrade check disabled with the unstable environment variable")
} else {
checkRedirect = func(req *http.Request, via []*http.Request) error {
if len(via) != 0 && strings.EqualFold(via[len(via)-1].URL.Scheme, "https") && !strings.EqualFold(req.URL.Scheme, "https") {
return errors.New("connection downgrade not allowed for URL: " + req.URL.String())
}
if len(via) >= 10 {
return errors.New("stopped after 10 redirects")
}
return nil
}
}
httpClient := &http.Client{
CheckRedirect: checkRedirect,
Transport: params.Options.Transport,
}
ctx, cancel := context.WithTimeout(ctx, defaults.DefaultIOTimeout)
defer cancel()
req, err := http.NewRequestWithContext(ctx, "GET", url, nil)
if err != nil {
return "", nil, trace.Wrap(err)
}
resp, err := httpClient.Do(req)
if err != nil {
return "", nil, trace.Wrap(err)
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return "", nil, trace.Errorf("unexpected status code %d", resp.StatusCode)
}
body, err := utils.ReadAtMost(resp.Body, teleport.MaxHTTPResponseSize)
if err != nil {
return "", nil, trace.Wrap(err)
}
rawEntityDescriptor = string(body)
if params.MFA {
log.DebugContext(ctx, "Successfully fetched MFA entity descriptor for connector")
} else {
log.DebugContext(ctx, "Successfully fetched entity descriptor for connector")
}
}
if rawEntityDescriptor == "" {
return "", nil, nil
}
entityDescriptor := &samltypes.EntityDescriptor{}
if err := xml.Unmarshal([]byte(rawEntityDescriptor), entityDescriptor); err != nil {
return "", nil, trace.BadParameter("failed to parse entity_descriptor XML: %s", err)
}
return rawEntityDescriptor, entityDescriptor, nil
}
// SAMLAssertionsToTraits converts saml assertions to traits
func SAMLAssertionsToTraits(assertions saml2.AssertionInfo) map[string][]string {
traits := make(map[string][]string, len(assertions.Values))
for _, assr := range assertions.Values {
vals := make([]string, 0, len(assr.Values))
for _, value := range assr.Values {
vals = append(vals, value.Value)
}
traits[assr.Name] = vals
}
return traits
}
// CheckSAMLEntityDescriptor checks if the entity descriptor XML is valid and has at least one valid certificate.
func CheckSAMLEntityDescriptor(entityDescriptor string) ([]*x509.Certificate, error) {
if entityDescriptor == "" {
return nil, nil
}
metadata := &samltypes.EntityDescriptor{}
if err := xml.Unmarshal([]byte(entityDescriptor), metadata); err != nil {
return nil, trace.Wrap(err, "failed to parse entity_descriptor")
}
if metadata.IDPSSODescriptor == nil {
return nil, nil
}
var roots []*x509.Certificate
for _, kd := range metadata.IDPSSODescriptor.KeyDescriptors {
for _, samlCert := range kd.KeyInfo.X509Data.X509Certificates {
// The certificate is base64 encoded and can be split into multiple lines.
// Each line can be padded with spaces/tabs, so we need to remove them first
// before decoding otherwise we'll get an error.
// We need to run this through strings.Fields to remove spaces/tabs
// from each line and then join them back with newlines.
// The last step isn't strictly necessary, but it makes payload more readable.
certData, err := base64.StdEncoding.DecodeString(strings.Join(strings.Fields(samlCert.Data), "\n"))
if err != nil {
return nil, trace.Wrap(err, "failed to decode certificate defined in entity_descriptor")
}
cert, err := x509.ParseCertificate(certData)
if err != nil {
return nil, trace.Wrap(err, "failed to parse certificate defined in entity_descriptor")
}
roots = append(roots, cert)
}
}
return roots, nil
}
// GetSAMLServiceProvider gets the SAMLConnector's service provider
func GetSAMLServiceProvider(sc types.SAMLConnector, clock clockwork.Clock) (*saml2.SAMLServiceProvider, error) {
roots, errEd := CheckSAMLEntityDescriptor(sc.GetEntityDescriptor())
if errEd != nil {
return nil, trace.Wrap(errEd)
}
certStore := dsig.MemoryX509CertificateStore{Roots: roots}
if sc.GetCert() != "" {
cert, err := tlsca.ParseCertificatePEM([]byte(sc.GetCert()))
if err != nil {
return nil, trace.Wrap(err, "failed to parse certificate defined in cert")
}
certStore.Roots = append(certStore.Roots, cert)
}
if len(certStore.Roots) == 0 {
return nil, trace.BadParameter("no identity provider certificate provided, either set certificate as a parameter or via entity_descriptor")
}
signingKeyPair := sc.GetSigningKeyPair()
encryptionKeyPair := sc.GetEncryptionKeyPair()
var keyStore *utils.KeyStore
var signingKeyStore *utils.KeyStore
var err error
// Due to some weird design choices with how gosaml2 keys are configured we have to do some trickery
// in order to default properly when SAML assertion encryption is turned off.
// Below are the different possible cases.
if encryptionKeyPair == nil {
// Case 1: Only the signing key pair is set. This means that SAML encryption is not expected
// and we therefore configure the main key that gets used for all operations as the signing key.
// This is done because gosaml2 mandates an encryption key even if not used.
slog.InfoContext(context.Background(), "No assertion_key_pair was detected, falling back to signing key for all SAML operations",
teleport.ComponentKey, teleport.ComponentSAML,
)
keyStore, err = utils.ParseKeyStorePEM(signingKeyPair.PrivateKey, signingKeyPair.Cert)
signingKeyStore = keyStore
if err != nil {
return nil, trace.Wrap(err, "failed to parse certificate or private key defined in signing_key_pair")
}
} else {
// Case 2: An encryption keypair is configured. This means that encrypted SAML responses are expected.
// Since gosaml2 always uses the main key for encryption, we set it to assertion_key_pair.
// To handle signing correctly, we now instead set the optional signing key in gosaml2 to signing_key_pair.
slog.InfoContext(context.Background(), "Detected assertion_key_pair and configured it to decrypt SAML responses",
teleport.ComponentKey, teleport.ComponentSAML,
)
keyStore, err = utils.ParseKeyStorePEM(encryptionKeyPair.PrivateKey, encryptionKeyPair.Cert)
if err != nil {
return nil, trace.Wrap(err, "failed to parse certificate or private key defined in assertion_key_pair")
}
signingKeyStore, err = utils.ParseKeyStorePEM(signingKeyPair.PrivateKey, signingKeyPair.Cert)
if err != nil {
return nil, trace.Wrap(err, "failed to parse certificate or private key defined in signing_key_pair")
}
}
sp := &saml2.SAMLServiceProvider{
IdentityProviderSSOURL: sc.GetSSO(),
IdentityProviderIssuer: sc.GetIssuer(),
ServiceProviderIssuer: sc.GetServiceProviderIssuer(),
AssertionConsumerServiceURL: sc.GetAssertionConsumerService(),
SignAuthnRequests: true,
SignAuthnRequestsCanonicalizer: dsig.MakeC14N11Canonicalizer(),
AudienceURI: sc.GetAudience(),
IDPCertificateStore: &certStore,
SPSigningKeyStore: signingKeyStore,
SPKeyStore: keyStore,
Clock: dsig.NewFakeClock(clock),
NameIdFormat: "urn:oasis:names:tc:SAML:1.1:nameid-format:unspecified",
ForceAuthn: sc.GetForceAuthn(),
}
// Provider specific settings for ADFS and JumpCloud. Specifically these
// providers do not support C14N11, which means a C14N10 canonicalizer has to
// be used.
switch sc.GetProvider() {
case teleport.ADFS, teleport.JumpCloud:
slog.DebugContext(context.Background(), "Setting ADFS/JumpCloud values", teleport.ComponentKey, teleport.ComponentSAML)
if sp.SignAuthnRequests {
sp.SignAuthnRequestsCanonicalizer = dsig.MakeC14N10ExclusiveCanonicalizerWithPrefixList(dsig.DefaultPrefix)
// At a minimum we require password protected transport.
sp.RequestedAuthnContext = &saml2.RequestedAuthnContext{
Comparison: "minimum",
Contexts: []string{"urn:oasis:names:tc:SAML:2.0:ac:classes:PasswordProtectedTransport"},
}
}
}
return sp, nil
}
// UnmarshalSAMLConnector unmarshals the SAMLConnector resource from JSON.
func UnmarshalSAMLConnector(bytes []byte, opts ...MarshalOption) (types.SAMLConnector, error) {
return UnmarshalSAMLConnectorWithValidationOptions(bytes, nil, opts...)
}
// UnmarshalSAMLConnectorWithValidationOptions unmarshals the SAMLConnector resource from JSON.
func UnmarshalSAMLConnectorWithValidationOptions(bytes []byte, validationOpts []types.SAMLConnectorValidationOption, opts ...MarshalOption) (types.SAMLConnector, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
err = utils.FastUnmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V2:
var c types.SAMLConnectorV2
if err := utils.FastUnmarshal(bytes, &c); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := ValidateSAMLConnector(&c, nil, validationOpts...); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
c.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
c.SetExpiry(cfg.Expires)
}
return &c, nil
}
return nil, trace.BadParameter("SAML connector resource version %v is not supported", h.Version)
}
// MarshalSAMLConnector marshals the SAMLConnector resource to JSON.
func MarshalSAMLConnector(samlConnector types.SAMLConnector, opts ...MarshalOption) ([]byte, error) {
if err := ValidateSAMLConnector(samlConnector, nil, types.SAMLConnectorValidationWithAttributesToRoles(true)); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch samlConnector := samlConnector.(type) {
case *types.SAMLConnectorV2:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, samlConnector))
default:
return nil, trace.BadParameter("unrecognized SAML connector version %T", samlConnector)
}
}
// FillSAMLSecretFieldsFromExistingConnector looks up the existing SAML connector and populates the secret fields if any are missing.
func FillSAMLSecretFieldsFromExistingConnector(ctx context.Context, connector types.SAMLConnector, sg SAMLConnectorGetter) error {
getExisting := sync.OnceValues(func() (types.SAMLConnector, error) {
return sg.GetSAMLConnectorWithValidationOptions(ctx, connector.GetName(), true, types.SAMLConnectorValidationFollowURLs(false))
})
if err := fillSAMLSigningKeyFromExisting(connector, getExisting); err != nil {
return trace.Wrap(err)
}
if err := fillSAMLOAuthClientSecretFromExisting(connector, getExisting); err != nil {
return trace.Wrap(err)
}
return nil
}
// fillSAMLSigningKeyFromExisting populates the signing key on the given connector if it's missing with the signing key
// from the connector returned by getExisting.
func fillSAMLSigningKeyFromExisting(connector types.SAMLConnector, getExisting samlConnectorGetter) error {
connectorSKP := connector.GetSigningKeyPair()
if connectorSKP == nil {
return nil
}
if connectorSKP.PrivateKey != "" {
return nil
}
existing, err := getExisting()
if err != nil {
return trace.Wrap(err)
}
keyPair := existing.GetSigningKeyPair()
if keyPair == nil {
return trace.BadParameter("the SAML connector has no signing key pair and none was provided. " + ErrMsgHowToFixMissingPrivateKey)
}
if keyPair.Cert != connectorSKP.Cert {
return trace.BadParameter("the SAML connector signing key cert does not match the existing one. " + ErrMsgHowToFixMissingPrivateKey)
}
connector.SetSigningKeyPair(keyPair)
return nil
}
// fillSAMLOAuthClientSecretFromExisting populates the client secret on the given connector if it's missing with the client secret
// from the connector returned by getExisting.
func fillSAMLOAuthClientSecretFromExisting(connector types.SAMLConnector, getExisting samlConnectorGetter) error {
connectorOAuthCreds := connector.GetOAuthClientCredentials()
if connectorOAuthCreds == nil {
return nil
}
if connectorOAuthCreds.ClientSecret != "" {
return nil
}
existing, err := getExisting()
if err != nil {
return trace.Wrap(err)
}
oauthCreds := existing.GetOAuthClientCredentials()
if oauthCreds == nil {
return trace.BadParameter("the existing SAML connector has no OAuth credentials. " + ErrMsgHowToFixMissingOAuthCreds)
}
if oauthCreds.ClientId != connectorOAuthCreds.ClientId {
return trace.BadParameter("the SAML connector client ID does not match the existing one. " + ErrMsgHowToFixMissingOAuthCreds)
}
connector.SetOAuthClientCredentials(oauthCreds)
return nil
}
/*
* 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"
"fmt"
"log/slog"
"net/url"
"slices"
"strings"
"github.com/crewjam/saml"
"github.com/crewjam/saml/samlsp"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// SAMLIdpServiceProviderGetter defines interface for fetching SAMLIdPServiceProvider resources.
type SAMLIdpServiceProviderGetter interface {
ListSAMLIdPServiceProviders(ctx context.Context, pageSize int, nextKey string) ([]types.SAMLIdPServiceProvider, string, error)
}
// SAMLIdPServiceProviders defines an interface for managing SAML IdP service providers.
type SAMLIdPServiceProviders interface {
SAMLIdpServiceProviderGetter
// GetSAMLIdPServiceProvider returns the specified SAML IdP service provider resources.
GetSAMLIdPServiceProvider(ctx context.Context, name string) (types.SAMLIdPServiceProvider, error)
// CreateSAMLIdPServiceProvider creates a new SAML IdP service provider resource.
CreateSAMLIdPServiceProvider(context.Context, types.SAMLIdPServiceProvider) error
// UpdateSAMLIdPServiceProvider updates an existing SAML IdP service provider resource.
UpdateSAMLIdPServiceProvider(context.Context, types.SAMLIdPServiceProvider) error
// DeleteSAMLIdPServiceProvider removes the specified SAML IdP service provider resource.
DeleteSAMLIdPServiceProvider(ctx context.Context, name string) error
// DeleteAllSAMLIdPServiceProviders removes all SAML IdP service providers.
DeleteAllSAMLIdPServiceProviders(context.Context) error
}
// MarshalSAMLIdPServiceProvider marshals the SAMLIdPServiceProvider resource to JSON.
func MarshalSAMLIdPServiceProvider(serviceProvider types.SAMLIdPServiceProvider, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch sp := serviceProvider.(type) {
case *types.SAMLIdPServiceProviderV1:
if err := sp.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, sp))
default:
return nil, trace.BadParameter("unsupported SAML IdP service provider resource %T", sp)
}
}
// UnmarshalSAMLIdPServiceProvider unmarshals SAMLIdPServiceProvider resource from JSON.
func UnmarshalSAMLIdPServiceProvider(data []byte, opts ...MarshalOption) (types.SAMLIdPServiceProvider, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing SAML IdP service provider data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var s types.SAMLIdPServiceProviderV1
if err := utils.FastUnmarshal(data, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
return &s, nil
}
return nil, trace.BadParameter("unsupported SAML IdP service provider resource version %q", h.Version)
}
// supportedACSBindings is the set of AssertionConsumerService bindings that teleport supports.
var supportedACSBindings = map[string]struct{}{
saml.HTTPPostBinding: {},
saml.HTTPRedirectBinding: {},
}
// ValidateAssertionConsumerService checks if a given assertion consumer service is usable by teleport. Note that
// it is permissible for a service provider to include acs endpoints that are not compatible with teleport, so long
// as at least one _is_ compatible.
func ValidateAssertionConsumerService(acs saml.IndexedEndpoint) error {
if _, ok := supportedACSBindings[acs.Binding]; !ok {
return trace.BadParameter("unsupported acs binding: %q", acs.Binding)
}
if acs.Location == "" {
return trace.BadParameter("acs location endpoint is missing or could not be decoded for %q binding", acs.Binding)
}
return trace.Wrap(validateAssertionConsumerServicesEndpoint(acs.Location))
}
// FilterSAMLEntityDescriptor performs a filter in place to remove unsupported and/or insecure fields from
// a saml entity descriptor. Specifically, it removes acs endpoints that are either of an unsupported kind,
// or are using a non-https endpoint. We perform filtering rather than outright rejection because it is generally
// expected that a service provider will successfully support a given ACS so long as they have at least one
// compatible binding.
func FilterSAMLEntityDescriptor(ed *saml.EntityDescriptor, quiet bool) error {
var originalCount int
var filteredCount int
for i := range ed.SPSSODescriptors {
filtered := slices.DeleteFunc(ed.SPSSODescriptors[i].AssertionConsumerServices, func(acs saml.IndexedEndpoint) bool {
if err := ValidateAssertionConsumerService(acs); err != nil {
if !quiet {
slog.WarnContext(context.Background(), "AssertionConsumerService binding for entity is invalid and will be ignored",
"entity_id", ed.EntityID,
"error", err,
)
}
return true
}
return false
})
originalCount += len(ed.SPSSODescriptors[i].AssertionConsumerServices)
filteredCount += len(filtered)
ed.SPSSODescriptors[i].AssertionConsumerServices = filtered
}
if filteredCount != originalCount {
return trace.BadParameter("Entity descriptor for entity %q contains unsupported AssertionConsumerService binding or location", ed.EntityID)
}
return nil
}
// invalidSAMLIdPACSURLChars contains low hanging HTML tag characters that are more
// commonly used in xss payload. This is not a comprehensive list but is only
// meant to increase the ost of xss payload.
const invalidSAMLIdPACSURLChars = `<>"!;`
// SAMLACSInputFilteringThreshold defines level of strictness for entity descriptor filtering.
type SAMLACSInputFilteringThreshold string
const (
// SAMLACSInputStrictFilter indicates ValidateAndFilterEntityDescriptor to return an error on
// any instance of unsupported ACS value.
SAMLACSInputStrictFilter SAMLACSInputFilteringThreshold = "SAMLACSInputStrictFilter"
// SAMLACSInputPermissiveFilter indicates ValidateAndFilterEntityDescriptor to ignore an error on
// any instance of unsupported ACS value.
SAMLACSInputPermissiveFilter SAMLACSInputFilteringThreshold = "SAMLACSInputPermissiveFilter"
)
// ValidateAndFilterEntityDescriptor validates entity id and ACS value. It specifically:
// - checks for a valid entity descriptor XML format.
// - checks for a matching entity ID field in both the entity_id field and entity ID contained in the value of
// entity_descriptor field.
// - performs filtering on the Assertion Consumer service (ACS) binding format or its location URL endpoint.
// filterThreshold dictates if ValidateAndFilterEntityDescriptor should return or ignore error on filtering result.
func ValidateAndFilterEntityDescriptor(sp types.SAMLIdPServiceProvider, filterThreshold SAMLACSInputFilteringThreshold) error {
edOriginal, err := samlsp.ParseMetadata([]byte(sp.GetEntityDescriptor()))
if err != nil {
return trace.BadParameter("invalid entity descriptor for SAML IdP Service Provider %q: %v", sp.GetEntityID(), err)
}
if edOriginal.EntityID != sp.GetEntityID() {
return trace.BadParameter("entity ID parsed from the entity descriptor does not match the entity ID in the SAML IdP service provider object")
}
if err := FilterSAMLEntityDescriptor(edOriginal, false /* quiet */); err != nil {
if filterThreshold == SAMLACSInputStrictFilter {
return trace.BadParameter("Entity descriptor for SAML IdP Service Provider %q contains unsupported ACS bindings: %v", sp.GetEntityID(), err)
}
}
return nil
}
// validateAssertionConsumerServicesEndpoint ensures that the Assertion Consumer Service location
// is a valid HTTPS endpoint.
func validateAssertionConsumerServicesEndpoint(acs string) error {
endpoint, err := url.Parse(acs)
switch {
case err != nil:
return trace.BadParameter("acs location endpoint %q could not be parsed: %v", acs, err)
case endpoint.Scheme != "https":
return trace.BadParameter("invalid scheme %q in acs location endpoint %q (must be 'https')", endpoint.Scheme, acs)
}
if strings.ContainsAny(acs, invalidSAMLIdPACSURLChars) {
return trace.BadParameter("acs location endpoint contains an unsupported character")
}
return nil
}
// ValidateSAMLIdPACSURLAndRelayStateInputs performs validation on SAML IdP Service Provider
// ACS URL and Relay State fields.
func ValidateSAMLIdPACSURLAndRelayStateInputs(sp types.SAMLIdPServiceProvider) error {
if sp.GetACSURL() != "" {
if err := validateAssertionConsumerServicesEndpoint(sp.GetACSURL()); err != nil {
return trace.Wrap(err)
}
}
if strings.ContainsAny(sp.GetRelayState(), invalidSAMLIdPACSURLChars) {
return trace.BadParameter("relay state contains an unsupported character")
}
return nil
}
// NewSAMLTestSPMetadata creates a new entity descriptor for tests.
func NewSAMLTestSPMetadata(entityID, acsURL string) string {
return fmt.Sprintf(samlTestSPMetadata, entityID, acsURL)
}
// samlTestSPMetadata mimics metadata format generated by saml.ServiceProvider.Metadata()
const samlTestSPMetadata = `<EntityDescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" validUntil="2023-12-09T23:43:58.16Z" entityID="%s">
<SPSSODescriptor xmlns="urn:oasis:names:tc:SAML:2.0:metadata" validUntil="2023-12-09T23:43:58.16Z" protocolSupportEnumeration="urn:oasis:names:tc:SAML:2.0:protocol" AuthnRequestsSigned="false" WantAssertionsSigned="true">
<NameIDFormat>urn:oasis:names:tc:SAML:1.1:nameid-format:unspecified</NameIDFormat>
<AssertionConsumerService Binding="urn:oasis:names:tc:SAML:2.0:bindings:HTTP-POST" Location="%s" index="1"></AssertionConsumerService>
</SPSSODescriptor>
</EntityDescriptor>
`
/*
* Teleport
* Copyright (C) 2025 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"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
scopedaccessv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/access/v1"
workloadidentityv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/wrappers"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/lib/scopes"
scopedaccess "github.com/gravitational/teleport/lib/scopes/access"
"github.com/gravitational/teleport/lib/scopes/pinning"
)
// ErrScopedIdentity is returned when a component intended for use only with unscoped identities receives a scoped
// identity. Methods that implement scoping support may check for this error and fallback to scoped authorization
// as appropriate.
var ErrScopedIdentity = &trace.AccessDeniedError{
Message: "scoped identities not supported",
}
// errUnenforcceableAssignment indicates that the role assignment is unenforceable due to scoping rules. This error
// is only emitted in contexts where it will be logged by the caller with the role's name and assignment scope.
var errUnenforcceableAssignment = &trace.AccessDeniedError{Message: "role's scoping does not permit enforcement as assigned"}
// errMissingAssignedRole indicates that the role assigned to a scoped identity was not found. This error is only
// emitted in contexts where it will be logged by the caller with the role's name and assignment scope.
var errMissingAssignedRole = &trace.AccessDeniedError{Message: "assigned role was not found"}
// scopedAccessCheckerBuilder is a helper that builds scoped access checkers.
type scopedAccessCheckerBuilder struct {
info *AccessInfo
localCluster string
reader ScopedRoleReader
}
// Check verifies that the builder was provided with all necessary parameters and that they are well-formed.
func (b *scopedAccessCheckerBuilder) Check() error {
if b.reader == nil {
return trace.BadParameter("cannot create scoped access checkers without a scoped role reader")
}
if b.localCluster == "" {
return trace.BadParameter("cannot create scoped access checkers without a local cluster name")
}
if b.info.ScopePin == nil {
return trace.BadParameter("cannot create scoped access checkers for unscoped identity")
}
if len(b.info.AllowedResourceAccessIDs) != 0 {
return trace.BadParameter("cannot create scoped access checkers for identity with active resource IDs")
}
if err := pinning.WeakValidate(b.info.ScopePin); err != nil {
return trace.Errorf("cannot create scoped access checkers: %w", err)
}
return nil
}
func (b *scopedAccessCheckerBuilder) newCheckerForRole(ctx context.Context, key pinning.RoleAssignment) (*ScopedAccessChecker, error) {
if key == (pinning.RoleAssignment{}) {
return b.newDefaultImplicitChecker(ctx), nil
}
if key.RoleKind != pinning.RoleKindUser {
return nil, trace.BadParameter("cannot build checker for non-user role kind %q (this is a bug)", key.RoleKind)
}
rsp, err := b.reader.GetScopedRole(ctx, scopedaccessv1.GetScopedRoleRequest_builder{
Name: key.RoleName,
Scope: key.RoleScope,
}.Build())
if err != nil {
if trace.IsNotFound(err) {
return nil, errMissingAssignedRole
}
return nil, trace.Wrap(err)
}
// verify that the role is enforceable at this enforcement point. if not, the assignment is
// skipped. this check is a critical part of the scopes security model and must always be
// performed prior to any enforcement logic related to a scoped role.
if !scopedaccess.RoleIsEnforceableAt(rsp.GetRole(), scopes.EnforcementPoint{
ScopeOfOrigin: key.ScopeOfOrigin,
ScopeOfEffect: key.ScopeOfEffect,
}) {
return nil, errUnenforcceableAssignment
}
// Convert the scoped role to a classic role using the scope of effect.
// The scope of effect determines which resources this role's privileges apply to.
role, err := scopedaccess.ScopedRoleToRole(rsp.GetRole(), key.ScopeOfEffect)
if err != nil {
return nil, trace.Wrap(err)
}
// TODO(fspmarshall/scopes): figure out how/when we want to support trait interpolation in scoped
// roles. When we do, that will likely need to be done here.
// Create an access checker with this single role. Single-role evaluation is a core principle
// of the scoped access model - the first role that permits access determines all parameters.
checker := newAccessChecker(b.info, b.localCluster, newScopedRoleSet(role))
return &ScopedAccessChecker{
scopeOfOrigin: key.ScopeOfOrigin,
scopeOfEffect: key.ScopeOfEffect,
role: rsp.GetRole(),
scopedCompatChecker: checker,
}, nil
}
// newScopedRoleSet builds the classic role set backing a scoped identity's compat checker. It exists
// instead of [NewRoleSet] because that appends the classic default implicit role, which grants
// secret-inclusive read; scoped identities get [newScopedImplicitRole] instead. Note that this must be
// used for *every* scoped checker, since the implicit role is appended to each single-role checker.
func newScopedRoleSet(roles ...types.Role) RoleSet {
return append(roles, newScopedImplicitRole())
}
// newDefaultImplicitChecker builds a scoped access checker representing the default implicit role. We rely on the privileges conferred
// by the default implicit role always being "assigned" at root as if they came from a root scoped role assignment. We achieve this by
// creating a fake scoped access checker that wraps an unscoped access checker holding only the scoped implicit role, which effectively
// simulates the presence of the default implicit role at root scope. Note that as functionality of scoped roles further diverges from
// unscoped roles, we may need to revisit this approach in favor of defining our own default implicit scoped role instead.
func (b *scopedAccessCheckerBuilder) newDefaultImplicitChecker(_ context.Context) *ScopedAccessChecker {
return &ScopedAccessChecker{
scopeOfOrigin: scopes.Root,
scopeOfEffect: scopes.Root,
scopedCompatChecker: newAccessChecker(b.info, b.localCluster, newScopedRoleSet()),
role: scopedaccessv1.ScopedRole_builder{
Metadata: headerv1.Metadata_builder{
Name: constants.DefaultImplicitRole,
}.Build(),
Scope: scopes.Root,
Spec: scopedaccessv1.ScopedRoleSpec_builder{
AssignableScopes: []string{scopes.Root},
}.Build(),
Version: types.V1,
}.Build(),
}
}
// ScopedAccessChecker performs access checks abstracting over scoped and unscoped identities.
//
// For scoped identities, each ScopedAccessChecker represents a single role assignment characterized by:
// - Scope of Origin: the scope from which the assignment originates (determines seniority)
// - Scope of Effect: the scope at which the role's privileges apply (determines applicability)
//
// In the scoped access model, the first role (in evaluation order) that permits access to a resource
// determines all subsequent access parameters. This differs from classic role evaluation where roles are
// aggregated and the most restrictive settings win.
//
// For unscoped identities, the full AccessChecker is wrapped directly and all method calls are delegated to it.
//
// ScopedAccessChecker instances should be obtained from ScopedAccessCheckerContext rather than constructed
// directly. The exception is NewScopedAccessCheckerFromUnscoped for adapting an unscoped AccessChecker.
type ScopedAccessChecker struct {
// scopeOfOrigin/scopeOfEffect are populated only for scoped identities; zero for unscoped.
scopeOfOrigin string
scopeOfEffect string
// role is the scoped role being evaluated, or nil for unscoped identities.
role *scopedaccessv1.ScopedRole
// scopedCompatChecker is a classic AccessChecker built from the scoped role via ScopedRoleToRole.
// Non-nil iff isScoped(). Used for checks that fall back to compat classic-role logic.
scopedCompatChecker AccessChecker
// unscopedChecker is the underlying unscoped AccessChecker.
// Non-nil iff !isScoped().
unscopedChecker AccessChecker
}
// NewScopedAccessCheckerFromUnscoped creates a ScopedAccessChecker wrapping an unscoped AccessChecker.
// This is used in code paths that accept *ScopedAccessChecker but operate on an unscoped identity.
func NewScopedAccessCheckerFromUnscoped(checker AccessChecker) *ScopedAccessChecker {
return &ScopedAccessChecker{unscopedChecker: checker}
}
// NewScopedAccessCheckerForSystemRole creates a ScopedAccessChecker for a single system role. Currently
// the checker masquerades as a scoped role checker but pulls all meaningful functionality from the provided
// unscoped access checker. This is the simplest way to achieve our desired effect of having scoped agent
// system roles act like scoped roles, but is somewhat brittle. It only works right now because we happen
// to still defer to the scopedCompatChecker for resource access checks. In the long run we may want to
// consider providing true scoped role representations of system roles, or more likely representing the
// system role presets in a format suitable for representation as a scoped or unscoped role.
// TODO(fspmarshall/scopes): revisit our scoped system role strateg as described above.
func NewScopedAccessCheckerForSystemRole(roleName string, checker AccessChecker) *ScopedAccessChecker {
return &ScopedAccessChecker{
scopeOfOrigin: scopes.Root,
scopeOfEffect: scopes.Root,
role: scopedaccessv1.ScopedRole_builder{
Metadata: headerv1.Metadata_builder{
Name: "system/" + roleName,
}.Build(),
Scope: scopes.Root,
Version: types.V1,
Spec: scopedaccessv1.ScopedRoleSpec_builder{
AssignableScopes: []string{scopes.Root},
}.Build(),
}.Build(),
scopedCompatChecker: checker,
}
}
// isScoped reports whether this checker operates on a scoped identity.
func (c *ScopedAccessChecker) isScoped() bool {
return c.role != nil
}
// SSH returns an SSH-specific access checker backed by this checker. All SSH-specific methods
// (logins, port forwarding, recording mode, idle timeout, etc.) live on [SSHAccessChecker].
func (c *ScopedAccessChecker) SSH() *SSHAccessChecker {
return &SSHAccessChecker{checker: c}
}
// Kube returns a kube-specific access checker backed by this checker. All kube-specific methods
// (users, groups, idle timeout, etc.) live on [KubeAccessChecker].
func (c *ScopedAccessChecker) Kube() *KubeAccessChecker {
return &KubeAccessChecker{checker: c}
}
// App returns an app-specific access checker backed by this checker. All app-specific methods
// live on [AppAccessChecker].
func (c *ScopedAccessChecker) App() *AppAccessChecker {
return &AppAccessChecker{checker: c}
}
// AccessInfo returns the AccessInfo that this access checker is based on.
func (c *ScopedAccessChecker) AccessInfo() *AccessInfo {
if !c.isScoped() {
return c.unscopedChecker.AccessInfo()
}
return c.scopedCompatChecker.AccessInfo()
}
// Traits returns the set of user traits.
func (c *ScopedAccessChecker) Traits() wrappers.Traits {
// there is no concept of scoped traits currently, and none is planned or would be feasible at least
// until we've fully migrated to PDP and deprecated certificate-based traits.
if !c.isScoped() {
return c.unscopedChecker.Traits()
}
return c.scopedCompatChecker.Traits()
}
// CheckAccessToRules verifies that *all* of a series of verbs are permitted for the specified resource.
func (c *ScopedAccessChecker) CheckAccessToRules(ctx RuleContext, resource string, verbs ...scopedaccess.Verb) error {
if !c.isScoped() {
return checkAccessToRulesImpl(c.unscopedChecker, ctx, resource, verbs...)
}
// XXX: the sanity of [NewScopedAccessCheckerForSystemRole] depends upon us continuing to defer to
// scopedCompatChecker for resource permission checks. Any revisiting of this strategy must take
// our scoped system role strategy into account.
return checkAccessToRulesImpl(c.scopedCompatChecker, ctx, resource, verbs...)
}
// CheckAccessToRemoteCluster checks access to a remote cluster.
func (c *ScopedAccessChecker) CheckAccessToRemoteCluster(cluster types.RemoteCluster) error {
if !c.isScoped() {
return c.unscopedChecker.CheckAccessToRemoteCluster(cluster)
}
// remote cluster access is never permitted for scoped identities.
// NOTE: it is unclear whether or not this method should even be implemented for the scoped access checker. it may be more
// sensible to force outer enforcement logic to grapple with the fact that a scoped checker does not support remote clusters
// at the type-level. this has been implemented experimentally to explore the pattern of having the scoped access checker
// implement methods that always deny for unsupported features.
return trace.AccessDenied("remote cluster access is not permitted for scoped identities")
}
// CheckAccessToWorkloadIdentity checks access to a workload identity resource by
// matching it against the WorkloadIdentityLabels granted by the checker's roles.
// This is the scoped equivalent of the label-based access check used by the
// unscoped issuance path.
func (c *ScopedAccessChecker) CheckAccessToWorkloadIdentity(wi *workloadidentityv1pb.WorkloadIdentity) error {
if !c.isScoped() {
return c.unscopedChecker.CheckAccess(types.Resource153ToResourceWithLabels(wi), AccessState{})
}
return c.scopedCompatChecker.CheckAccess(types.Resource153ToResourceWithLabels(wi), AccessState{})
}
// AdjustSessionTTL will reduce the requested ttl to the lowest max allowed TTL for this role set.
func (c *ScopedAccessChecker) AdjustSessionTTL(ttl time.Duration) time.Duration {
// the naive implementation of this method for scopes may have problematic interactions with
// cert parameter generation. see ../scopes/access/compat.go for more detailed discussion.
if !c.isScoped() {
return c.unscopedChecker.AdjustSessionTTL(ttl)
}
return c.scopedCompatChecker.AdjustSessionTTL(ttl)
}
// PrivateKeyPolicy returns the enforced private key policy, or the provided default, whichever is stricter.
func (c *ScopedAccessChecker) PrivateKeyPolicy(defaultPolicy keys.PrivateKeyPolicy) (keys.PrivateKeyPolicy, error) {
// the naive implementation of this method for scopes may have problematic interactions with
// cert parameter generation. see ../scopes/access/compat.go for more detailed discussion.
if !c.isScoped() {
return c.unscopedChecker.PrivateKeyPolicy(defaultPolicy)
}
return c.scopedCompatChecker.PrivateKeyPolicy(defaultPolicy)
}
// PinSourceIP returns whether source IP pinning is enforced.
func (c *ScopedAccessChecker) PinSourceIP() bool {
// the naive implementation of this method for scopes may have problematic interactions with
// cert parameter generation. see ../scopes/access/compat.go for more detailed discussion.
if !c.isScoped() {
return c.unscopedChecker.PinSourceIP()
}
return c.scopedCompatChecker.PinSourceIP()
}
// LockingMode returns the locking mode to apply.
func (c *ScopedAccessChecker) LockingMode(defaultMode constants.LockingMode) constants.LockingMode {
// the naive implementation of this method for scopes may have problematic interactions with
// cert parameter generation. see ../scopes/access/compat.go for more detailed discussion.
if !c.isScoped() {
return c.unscopedChecker.LockingMode(defaultMode)
}
return c.scopedCompatChecker.LockingMode(defaultMode)
}
// DelegationSessionID returns the ID of the current Delegation Session.
func (c *ScopedAccessChecker) DelegationSessionID() string {
if !c.isScoped() {
return c.unscopedChecker.DelegationSessionID()
}
return c.scopedCompatChecker.DelegationSessionID()
}
// checkAccessToRulesImpl verifies that *all* of a series of verbs are permitted for the specified resource. This
// function differs from AccessChecker.CheckAccessToRule in that it does not support advanced context-based features
// or namespacing, and accepts a set of verbs all of which must evaluate to allow for the check to succeed.
func checkAccessToRulesImpl(checker AccessChecker, ctx RuleContext, resource string, verbs ...scopedaccess.Verb) error {
if len(verbs) == 0 {
return trace.BadParameter("malformed rule check for %q, no verbs provided (this is a bug)", resource)
}
for _, verb := range verbs {
classicVerb, err := verb.ClassicVerb()
if err != nil {
return trace.Wrap(err, "malformed rule check for %q", resource)
}
if err := checker.CheckAccessToRule(ctx, apidefaults.Namespace, resource, classicVerb); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// checkMaybeHasAccessToRulesImpl returns an error if the checker definitely does not have access to the provided rules.
func checkMaybeHasAccessToRulesImpl(checker AccessChecker, ctx RuleContext, resource string, verbs ...scopedaccess.Verb) error {
if len(verbs) == 0 {
return trace.BadParameter("malformed maybe has access to rule check for %q, no verbs provided (this is a bug)", resource)
}
for _, verb := range verbs {
classicVerb, err := verb.ClassicVerb()
if err != nil {
return trace.Wrap(err, "malformed maybe has access to rule check for %q", resource)
}
if err := checker.GuessIfAccessIsPossible(ctx, apidefaults.Namespace, resource, classicVerb); err != nil {
return trace.Wrap(err)
}
}
return nil
}
/*
* Teleport
* Copyright (C) 2026 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"
"log/slog"
"github.com/gravitational/trace"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/wrappers"
dtauthz "github.com/gravitational/teleport/lib/devicetrust/authz"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/scopes"
scopedaccess "github.com/gravitational/teleport/lib/scopes/access"
"github.com/gravitational/teleport/lib/scopes/pinning"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils/once"
)
// ScopedAccessCheckerContext is the top-level access checker state, abstracting over scoped and unscoped
// identities. For scoped identities it builds and caches per-role checkers based on the scope pin and role
// assignments or system roles. For unscoped identities it wraps a standard AccessChecker.
//
// User-vs-agent differences are fully captured at construction time in the checkersAtPoint closure.
// Once constructed, this type is uniform across both scoped identity kinds.
type ScopedAccessCheckerContext struct {
// pin is the scope pin for this identity. Non-nil iff isScoped().
pin *scopesv1.Pin
// traits are the user traits for this context. Nil for agent pin identities.
traits wrappers.Traits
// resolveRef resolves a RoleAssignment to a ScopedAccessChecker. The zero-value RoleAssignment is
// the preamble sentinel: user pins return the default implicit role checker, agent pins return nil.
// Non-nil iff isScoped().
resolveRef func(ctx context.Context, ref pinning.RoleAssignment) (*ScopedAccessChecker, error)
// enumerateAll enumerates checkers across all role assignments, for cert parameter aggregation.
// Non-nil only for user pins; nil for agent pins (cert param aggregation is not supported).
enumerateAll func(ctx context.Context) stream.Stream[*ScopedAccessChecker]
// unscopedChecker wraps a standard AccessChecker for unscoped identities.
// Non-nil iff !isScoped().
unscopedChecker AccessChecker
}
// NewScopedAccessCheckerContext builds a ScopedAccessCheckerContext for a scoped user identity.
func NewScopedAccessCheckerContext(ctx context.Context, info *AccessInfo, localCluster string, reader ScopedRoleReader) (*ScopedAccessCheckerContext, error) {
builder := scopedAccessCheckerBuilder{
info: info,
localCluster: localCluster,
reader: reader,
}
if err := builder.Check(); err != nil {
return nil, trace.Wrap(err)
}
pin := info.ScopePin
if pin.GetKind() != scopesv1.PinKind_PIN_KIND_USER {
return nil, trace.BadParameter("cannot create user pin checker context for pin of kind %v", pin.GetKind())
}
cachedCheckerForRole, _ := once.KeyedValue(builder.newCheckerForRole)
resolveRef := func(ctx context.Context, ref pinning.RoleAssignment) (*ScopedAccessChecker, error) {
if ref == (pinning.RoleAssignment{}) {
// preamble sentinel: return the default implicit role checker
return cachedCheckerForRole(ctx, pinning.RoleAssignment{})
}
return cachedCheckerForRole(ctx, ref)
}
enumerateAll := func(ctx context.Context) stream.Stream[*ScopedAccessChecker] {
return func(yield func(*ScopedAccessChecker, error) bool) {
var yielded int
var lastErr error
for assignment := range pinning.EnumerateAllAssignments(pin) {
checker, err := cachedCheckerForRole(ctx, assignment)
if err != nil {
slog.WarnContext(ctx, "skipping role evaluation due to error", "role_name", assignment.RoleName, "scope_of_origin", assignment.ScopeOfOrigin, "scope_of_effect", assignment.ScopeOfEffect, "error", err)
lastErr = err
continue
}
if !yield(checker, nil) {
return
}
yielded++
}
if yielded == 0 && lastErr != nil {
yield(nil, lastErr)
}
}
}
return &ScopedAccessCheckerContext{
pin: pin,
traits: info.Traits,
resolveRef: resolveRef,
enumerateAll: enumerateAll,
}, nil
}
// NewScopedAccessCheckerContextFromUnscoped builds a ScopedAccessCheckerContext wrapping an unscoped AccessChecker.
func NewScopedAccessCheckerContextFromUnscoped(checker AccessChecker) *ScopedAccessCheckerContext {
return &ScopedAccessCheckerContext{unscopedChecker: checker}
}
// NewScopedAccessCheckerContextForAgentPin builds a ScopedAccessCheckerContext for a scoped agent identity.
// Each entry in checkersByRole maps a system role name to its [ScopedAccessChecker], which must have been
// built via [NewScopedAccessCheckerForSystemRole].
//
// System role checkers are resolved at the (root, root) enforcement point, reflecting that system role
// permissions are treated as assigned at root scope. Note that this is not necessarily a permanent design
// choice. Future iterations may choose to represent system role permissions as a combination of root assigned
// permissions and agent scope assigned permissions. There are pros and cons to either appoach, but we've opted
// to start with the simpler model for now.
func NewScopedAccessCheckerContextForAgentPin(pin *scopesv1.Pin, checkersByRole map[string]*ScopedAccessChecker) (*ScopedAccessCheckerContext, error) {
if pin == nil {
return nil, trace.BadParameter("cannot create scoped access checker context without agent pin")
}
if pin.GetKind() != scopesv1.PinKind_PIN_KIND_AGENT {
return nil, trace.BadParameter("cannot create scoped access checker context for unexpected pin kind %v, expected %v", pin.GetKind(), scopesv1.PinKind_PIN_KIND_AGENT)
}
if len(checkersByRole) == 0 {
return nil, trace.BadParameter("cannot create scoped access checker context without any system role checkers")
}
if err := pinning.WeakValidate(pin); err != nil {
return nil, trace.Wrap(err)
}
resolveRef := func(ctx context.Context, ref pinning.RoleAssignment) (*ScopedAccessChecker, error) {
if ref == (pinning.RoleAssignment{}) {
// preamble sentinel: agent pins have no implicit role preamble
return nil, nil
}
if ref.RoleKind != pinning.RoleKindSystem {
return nil, trace.BadParameter("agent pin resolver received non-system role kind %q (this is a bug)", ref.RoleKind)
}
checker, ok := checkersByRole[ref.RoleName]
if !ok {
return nil, trace.BadParameter("no checker found for system role %q", ref.RoleName)
}
return checker, nil
}
return &ScopedAccessCheckerContext{
pin: pin,
resolveRef: resolveRef,
}, nil
}
// isScoped reports whether this context operates on a scoped identity.
func (c *ScopedAccessCheckerContext) isScoped() bool {
return c.unscopedChecker == nil
}
// ScopePin returns the scope pin for the identity, if the identity is scoped (user or agent).
// Returns (nil, false) for unscoped identities.
func (c *ScopedAccessCheckerContext) ScopePin() (*scopesv1.Pin, bool) {
return c.pin, c.pin != nil
}
// CheckersForResourceScope returns a stream of ScopedAccessCheckers in evaluation order for the given resource
// scope. For scoped identities, this enforces pin compliance and yields per-role checkers ordered by scope of
// origin (ancestral to descendant) then scope of effect (descendant to ancestral). For unscoped identities,
// yields a single checker wrapping the full unscoped context.
//
// This is the mechanism that *must* be used for getting checkers when checking access to a resource.
//
// Callers may pass an empty string ("") as scope to indicate an unscoped resource. This will not be treated as a
// root scope resource - i.e. identities with privileges assigned in the root scope will not be able to access the
// resource.
func (c *ScopedAccessCheckerContext) CheckersForResourceScope(ctx context.Context, scope string) stream.Stream[*ScopedAccessChecker] {
if !c.isScoped() {
return stream.Once(NewScopedAccessCheckerFromUnscoped(c.unscopedChecker))
}
const enforcePinTrue = true
return c.checkersForResourceScope(ctx, scope, enforcePinTrue)
}
// riskyUnpinnedCheckersForResourceScope is equivalent to CheckersForResourceScope except that it bypasses
// enforcement of the pinning scope. This is a risky operation that should only be used for certain APIs that
// make an exception to pinning exclusion rules (e.g. allowing read operations for resources at a parent scope).
func (c *ScopedAccessCheckerContext) riskyUnpinnedCheckersForResourceScope(ctx context.Context, scope string) stream.Stream[*ScopedAccessChecker] {
if !c.isScoped() {
return stream.Once(NewScopedAccessCheckerFromUnscoped(c.unscopedChecker))
}
const enforcePinFalse = false
return c.checkersForResourceScope(ctx, scope, enforcePinFalse)
}
func (c *ScopedAccessCheckerContext) checkersForResourceScope(ctx context.Context, scope string, enforcePin bool) stream.Stream[*ScopedAccessChecker] {
return func(yield func(*ScopedAccessChecker, error) bool) {
// deny immediately if the resource scope is not subject to the pinned scope. note that this denial isn't just an
// optimization, we have to perform this check separately from whatever access checks are performed by particular
// checkers. This is vital as the pin scope itself may deny access to a resource that would be permitted by any
// particular role. For example, if a user has a scoped role assigned at /foo which grants access to all ssh
// nodes, but they are pinned to scope /foo/bar, even if a role at /foo permits access, the pin restricts
// access to only resources subject to /foo/bar.
if enforcePin {
if !pinning.PinAppliesToResourceScope(c.pin, scope) {
yield(nil, trace.AccessDenied(
"a resource in scope %q can't be manipulated from a session pinned to scope %q",
scope, c.pin.GetScope(),
))
return
}
} else if !pinning.PinCompatibleWithPolicyScope(c.pin, scope) {
// Unpinned reads should still not allow orthogonal reads.
yield(nil, trace.AccessDenied(
"a resource in scope %q can't be manipulated from a session pinned to orthogonal scope %q",
scope, c.pin.GetScope(),
))
return
}
var successfullyResolved int
var lastErr error
// resolve and yield preamble checker if one exists. this step is necessary for user identities in order to
// ensure that default implicit role permissions are always evaluated first and always priority equivalent
// to a root scope of origin. agent identities do not require this step as they always have checkers with
// a root scope of origin. Note that the preamble checker does not count toward successfullyResolved.
if preambleChecker, err := c.resolveRef(ctx, pinning.RoleAssignment{}); err != nil {
slog.WarnContext(ctx, "skipping default implicit role evaluation due to error", "error", err)
lastErr = err
} else if preambleChecker != nil {
if !yield(preambleChecker, nil) {
return
}
}
// iterate through the ordered enforcement points for this resource scope. policy evaluation by scope is ordered first by
// Scope of Origin (ancestral to descendant) and then by Scope of Effect (descendant to ancestral within each origin).
// We proceed through each permutation in order, evaluating any roles assigned at that specific point.
for point := range scopes.EnforcementPointsForResourceScope(scope) {
for ref := range pinning.GetRolesAtEnforcementPoint(c.pin, point) {
checker, err := c.resolveRef(ctx, ref)
if err != nil {
// in classic teleport access checking skipping a role would be unacceptable due to side effects and deny rules. the scoped model
// however relies on cross-role isolation and explicitly allows omission of roles.
slog.WarnContext(ctx, "skipping role evaluation due to error",
"role_name", ref.RoleName,
"scope_of_origin", ref.ScopeOfOrigin,
"scope_of_effect", ref.ScopeOfEffect,
"error", err)
lastErr = err
continue
}
if !yield(checker, nil) {
return
}
successfullyResolved++
}
}
if successfullyResolved == 0 && lastErr != nil {
// if we didn't successfully build any assignment-derived checkers and encountered errors, return the last error encountered
// as it may be indicative of some kind of systemic failure rather than a problem with a specific assignment.
yield(nil, lastErr)
}
}
}
// riskyEnumerateScopedCheckers returns a stream of all possible scoped access checkers for the identity,
// enumerating every role assignment in the pin's assignment tree. The order is undefined and must not be
// relied upon for access control decisions. This method panics if called on an unscoped context or an agent
// pin context — it is only meaningful for scoped user identities.
//
// Note that use of this method should be treated with extreme caution. Accidental misuse could easily
// result in a scope isolation violation.
func (c *ScopedAccessCheckerContext) riskyEnumerateScopedCheckers(ctx context.Context) stream.Stream[*ScopedAccessChecker] {
if !c.isScoped() {
panic("riskyEnumerateScopedCheckers called on an unscoped access checker context (this is a bug)")
}
if c.enumerateAll == nil {
panic("riskyEnumerateScopedCheckers called on an agent pin context (this is a bug)")
}
return c.enumerateAll(ctx)
}
// ResolveScopeFilter returns the scope filter that should be used for a list/read/watch request, applying
// identity-derived defaulting when the caller did not specify an explicit filter. An empty (nil or
// MODE_UNSPECIFIED) filter is replaced with the safe default for this identity: MODE_EXACT at the pinned
// scope for scoped identities, or MODE_UNSCOPED for unscoped identities. A filter with an explicit mode is
// returned unchanged.
//
// This is the single source of truth for scope-filter defaulting. API handlers should pass caller-provided
// filters through this method so that an omitted filter is consistently and safely defaulted, minimizing
// per-API defaulting logic. Note that this method only handles defaulting; callers are still responsible for
// validating (see [scopes.ValidateFilter]) and authorizing the resulting filter.
//
// There is one subtle exception to this default behavior. The event system treats requests to watch certain
// *always unscoped* kinds as having a default UNSCOPED filter, even for scoped callers. This behavior is only
// present for unscoped kinds that have an unscoped read exception in place (e.g. cert authorities).
func (c *ScopedAccessCheckerContext) ResolveScopeFilter(filter *scopesv1.Filter) *scopesv1.Filter {
if filter.GetMode() != scopesv1.Mode_MODE_UNSPECIFIED {
// the caller specified an explicit filter; return it unchanged.
return filter
}
if !c.isScoped() {
// unscoped callers default to matching unscoped resources only.
return scopesv1.Filter_builder{
Mode: scopesv1.Mode_MODE_UNSCOPED,
}.Build()
}
// scoped callers default to the minimal/safe MODE_EXACT at their pinned scope.
return scopesv1.Filter_builder{
Scope: c.pin.GetScope(),
Mode: scopesv1.Mode_MODE_EXACT,
}.Build()
}
// CheckMaybeHasAccessToRules returns an error if the context definitely does not have access to the provided
// rules. For scoped identities, always returns nil — the scoped access model evaluates permissions per-resource.
func (c *ScopedAccessCheckerContext) CheckMaybeHasAccessToRules(ctx RuleContext, resource string, verbs ...scopedaccess.Verb) error {
if !c.isScoped() {
return checkMaybeHasAccessToRulesImpl(c.unscopedChecker, ctx, resource, verbs...)
}
return nil
}
// Decision calls fn against each checker in the resource scope evaluation order until one of three
// conditions is met: (1) fn succeeds, (2) fn returns an explicitly denied error, or (3) all checkers
// have been exhausted (implicit deny).
//
// Unscoped contexts resolve to exactly one checker, an error returned in this case is unmodified in order to
// surface special unscoped requirements like trusted device or session MFA that only an unscoped checker can produce
// right now.
func (c *ScopedAccessCheckerContext) Decision(ctx context.Context, scope string, fn func(*ScopedAccessChecker) error) error {
return c.decision(c.CheckersForResourceScope(ctx, scope), fn)
}
func (c *ScopedAccessCheckerContext) decision(checkers stream.Stream[*ScopedAccessChecker], fn func(*ScopedAccessChecker) error) error {
for checker, err := range checkers {
if err != nil {
return trace.Wrap(err)
}
err = fn(checker)
switch {
case err == nil:
return nil
case !c.isScoped():
// This surfaces requirements like trusted device and session MFA
// that only the unscoped checker can produce right now.
// Rather than masking them into an explicit deny,
// return errors here that unscoped checkers return.
return trace.Wrap(err)
case IsAccessExplicitlyDenied(err):
return trace.Wrap(err)
default:
// implicit deny, continue to the next check
continue
}
}
return trace.AccessDenied("access denied (decision)")
}
// AccessStateFromSSHIdentity builds an AccessState from an SSH identity, abstracting over scoped and
// unscoped access state construction.
func (c *ScopedAccessCheckerContext) AccessStateFromSSHIdentity(ctx context.Context, ident *sshca.Identity, authPrefGetter AuthPreferenceGetter) (AccessState, error) {
if !c.isScoped() {
return AccessStateFromSSHIdentity(ctx, ident, c.unscopedChecker, authPrefGetter)
}
authPref, err := authPrefGetter.GetAuthPreference(ctx)
if err != nil {
return AccessState{}, trace.Wrap(err)
}
if authPref.GetRequireMFAType().IsSessionMFARequired() {
// TODO(fspmarshall/scopes): implement scoped MFA
// NOTE: this will require additional refactoring of relevant access-checking logic. currently, we often
// check MFA requirements *before* we determine access to the underlying resource, but a scoped MFA model
// will need to first determine the scope of access *before* we can determine whether MFA is required for that scope.
return AccessState{}, trace.AccessDenied("cannot perform scoped access when cluster-level MFA is required (scoped MFA is not implemented)")
}
return AccessState{
// MFA state is hard-coded here because scoped roles do not support MFA yet, and the above check should reject
// cases where cluster-level config would obligate MFA.
MFARequired: MFARequiredNever,
MFAVerified: false,
EnableDeviceVerification: true,
DeviceVerified: dtauthz.IsSSHDeviceVerified(ident),
IsBot: ident.IsBot(),
}, nil
}
// AccessStateFromTLSIdentity builds an AccessState from an TLS identity, abstracting over scoped and
// unscoped access state construction.
func (c *ScopedAccessCheckerContext) AccessStateFromTLSIdentity(ctx context.Context, ident *tlsca.Identity, authPrefGetter AuthPreferenceGetter) (AccessState, error) {
if !c.isScoped() {
return AccessStateFromTLSIdentity(ctx, ident, c.unscopedChecker, authPrefGetter)
}
authPref, err := authPrefGetter.GetAuthPreference(ctx)
if err != nil {
return AccessState{}, trace.Wrap(err)
}
if authPref.GetRequireMFAType().IsSessionMFARequired() {
// TODO(fspmarshall/scopes): implement scoped MFA
// NOTE: this will require additional refactoring of relevant access-checking logic. currently, we often
// check MFA requirements *before* we determine access to the underlying resource, but a scoped MFA model
// will need to first determine the scope of access *before* we can determine whether MFA is required for that scope.
return AccessState{}, trace.AccessDenied("cannot perform scoped access when cluster-level MFA is required (scoped MFA is not implemented)")
}
return AccessState{
// MFA state is hard-coded here because scoped roles do not support MFA yet, and the above check should reject
// cases where cluster-level config would obligate MFA.
MFARequired: MFARequiredNever,
MFAVerified: false,
EnableDeviceVerification: true,
DeviceVerified: dtauthz.IsTLSDeviceVerified(&ident.DeviceExtensions),
IsBot: ident.IsBot(),
}, nil
}
// Traits returns the user traits for this context. Agent pin identities have no traits.
func (c *ScopedAccessCheckerContext) Traits() wrappers.Traits {
if !c.isScoped() {
return c.unscopedChecker.Traits()
}
return c.traits
}
// CertParams returns a sub-context for resolving certificate parameters during certificate generation.
// This should not be used outside of certificate generation logic.
func (c *ScopedAccessCheckerContext) CertParams() *CertificateParameterContext {
return &CertificateParameterContext{ctx: c}
}
// RiskyAuthorizeUnpinnedRead authorizes a read-only access check that bypasses
// enforcement of the identity's pinned scope. This must only be used for
// specific APIs that make an exception to pinning exclusion rules (e.g.
// allowing read operations for resources at a parent scope). To avoid misuse,
// a specific [UnpinnedReadAuthorization] must be provided that will encode the
// effective scope of the access check and the allowed verbs.
func (c *ScopedAccessCheckerContext) RiskyAuthorizeUnpinnedRead(
ctx context.Context,
authz UnpinnedReadAuthorization,
ruleCtx RuleContext,
) error {
if err := authz.check(); err != nil {
return trace.Wrap(err, "invalid unpinned read authorization")
}
return c.decision(
c.riskyUnpinnedCheckersForResourceScope(ctx, authz.resourceScope),
func(checker *ScopedAccessChecker) error {
return checker.CheckAccessToRules(ruleCtx, authz.kind, authz.verbs...)
},
)
}
// RiskyAuthorizeUnpinnedReadWithScope extends [RiskyAuthorizeUnpinnedRead].
// It authorizes a read-only access check that bypasses
// enforcement of the identity's pinned scope, but ensures that the
// resource scope is related to the identity's pinned scope.
// This must only be used for specific APIs that make an exception to pinning exclusion rules (e.g.
// allowing read operations for resources at a parent scope). To avoid misuse,
// a specific [UnpinnedReadAuthorization] must be provided that will encode the
// effective scope of the access check and the allowed verbs. The scope provided
// must not be empty, and will be used to determine the enforcement.
func (c *ScopedAccessCheckerContext) RiskyAuthorizeUnpinnedReadWithScope(
ctx context.Context,
authz UnpinnedReadAuthorization,
ruleCtx RuleContext,
resourceScope string,
) error {
authz.resourceScope = resourceScope
return c.RiskyAuthorizeUnpinnedRead(ctx, authz, ruleCtx)
}
// RiskyAuthorizeUnpinnedEmitEvent authorizes a create-only access check that bypasses
// enforcement of the identity's pinned scope specifically for emitting audit events.
// It is special-cased to avoid misuse of unpinned writes.
func (c *ScopedAccessCheckerContext) RiskyAuthorizeUnpinnedEmitEvent(
ctx context.Context,
ruleCtx RuleContext,
) error {
if pin, ok := c.ScopePin(); ok {
if pin.GetKind() != scopesv1.PinKind_PIN_KIND_AGENT {
return trace.AccessDenied("unpinned authorization for audit event emission is only supported for agent pins")
}
}
return c.decision(
c.riskyUnpinnedCheckersForResourceScope(ctx, scopes.Root),
func(checker *ScopedAccessChecker) error {
return checker.CheckAccessToRules(ruleCtx, types.KindEvent, scopedaccess.Create)
},
)
}
// RiskyAuthorizeUnpinnedWriteEvent authorizes a write-only access check that bypasses
// enforcement of the identity's pinned scope specifically for creating and updating audit events.
// It is special-cased to avoid misuse of unpinned writes.
func (c *ScopedAccessCheckerContext) RiskyAuthorizeUnpinnedWriteEvent(
ctx context.Context,
ruleCtx RuleContext,
) error {
if pin, ok := c.ScopePin(); ok {
if pin.GetKind() != scopesv1.PinKind_PIN_KIND_AGENT {
return trace.AccessDenied("unpinned authorization for audit event emission is only supported for agent pins")
}
}
return c.decision(
c.riskyUnpinnedCheckersForResourceScope(ctx, scopes.Root),
func(checker *ScopedAccessChecker) error {
return checker.CheckAccessToRules(ruleCtx, types.KindEvent, scopedaccess.Create, scopedaccess.Update)
},
)
}
// UnpinnedReadAuthorization is a special authorization to complete an unscoped
// read-only access check. This is meant to be used for access checks on
// typically cluster-wide resources that need to be readable by identities with
// a pinned scope.
type UnpinnedReadAuthorization struct {
resourceScope string
kind string
verbs []scopedaccess.Verb
}
func (a UnpinnedReadAuthorization) check() error {
switch {
case a.kind == "":
return trace.BadParameter("missing kind")
case len(a.verbs) == 0:
return trace.BadParameter("missing verbs")
}
for _, verb := range a.verbs {
switch verb {
case scopedaccess.List, scopedaccess.Read:
default:
// note that Secrets verb in particular is not allowed here, unpinned reads should
// never include secrets.
return trace.BadParameter("invalid verb for unpinned read authorization: %q", verb)
}
}
if err := scopes.WeakValidate(a.resourceScope); err != nil {
return trace.Wrap(err, "invalid resourceScope")
}
return nil
}
var (
// UnpinnedReadCertAuthority is a special authorization to complete an
// unscoped access check to read a cert authority without secrets.
UnpinnedReadCertAuthority = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindCertAuthority,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadCertAuthorities is a special authorization to complete an
// unscoped access check to list and read a cert authorities without secrets.
UnpinnedReadCertAuthorities = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindCertAuthority,
verbs: []scopedaccess.Verb{scopedaccess.List, scopedaccess.Read},
}
// UnpinnedReadAuthServers is a special authorization to complete an
// unscoped access check to list and read auth server resources.
UnpinnedReadAuthServers = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindAuthServer,
verbs: []scopedaccess.Verb{scopedaccess.List, scopedaccess.Read},
}
// UnpinnedReadProxies is a special authorization to complete an
// unscoped access check to list and read proxy resources.
UnpinnedReadProxies = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindProxy,
verbs: []scopedaccess.Verb{scopedaccess.List, scopedaccess.Read},
}
// UnpinnedReadAuthPreference is a special authorization to complete an
// unscoped access check to read a cluster auth preference.
UnpinnedReadAuthPreference = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindClusterAuthPreference,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadVnetConfig is a special authorization to complete an
// unscoped access check to read a cluster VNet config.
UnpinnedReadVnetConfig = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindVnetConfig,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadSPIFFEFederation is a special authorization to complete an
// unscoped access check to read a SPIFFE federation.
UnpinnedReadSPIFFEFederation = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindSPIFFEFederation,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadSPIFFEFederations is a special authorization to complete an
// unscoped access check to list and read SPIFFE federations.
UnpinnedReadSPIFFEFederations = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindSPIFFEFederation,
verbs: []scopedaccess.Verb{scopedaccess.List, scopedaccess.Read},
}
// UnpinnedReadClusterNetworkingConfig is a special authorization to complete an
// unscoped access check to read a cluster networking config.
UnpinnedReadClusterNetworkingConfig = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindClusterNetworkingConfig,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadClusterName is a special authorization to complete an
// unscoped access check to read a cluster name.
UnpinnedReadClusterName = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindClusterName,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadSessionRecordingConfig is a special authorization to complete an
// unscoped access check to read a cluster session recording config.
UnpinnedReadSessionRecordingConfig = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindSessionRecordingConfig,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadScopedRole is a special authorization to complete an
// unscoped access check to read a scoped role.
UnpinnedReadScopedRole = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: scopedaccess.KindScopedRole,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadUser is a special authorization to complete a unscoped access check
// to read a user.
UnpinnedReadUser = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindUser,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadRole is a special authorization to complete an unscoped access check
// to read a role.
UnpinnedReadRole = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindRole,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadAndListLock is a special authorization to complete an unscoped access check
// to read a lock.
UnpinnedReadAndListLock = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindLock,
verbs: []scopedaccess.Verb{scopedaccess.List, scopedaccess.Read},
}
// UnpinnedReadLock is a special authorization to complete an unscoped access check
// to read a lock.
UnpinnedReadLock = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindLock,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
// UnpinnedReadClusterAuditConfig is a special authorization to complete an unscoped access check
// to read a cluster audit config.
UnpinnedReadClusterAuditConfig = UnpinnedReadAuthorization{
resourceScope: scopes.Root,
kind: types.KindClusterAuditConfig,
verbs: []scopedaccess.Verb{scopedaccess.Read},
}
)
/*
* 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/trace"
"github.com/gravitational/teleport/api/types/secreports"
"github.com/gravitational/teleport/lib/utils"
)
// SecurityAuditQueryGetter is the interface for audit query getters.
type SecurityAuditQueryGetter interface {
// GetSecurityAuditQuery returns an audit query.
GetSecurityAuditQuery(ctx context.Context, name string) (*secreports.AuditQuery, error)
// GetSecurityAuditQueries returns all audit queries.
GetSecurityAuditQueries(context.Context) ([]*secreports.AuditQuery, error)
// ListSecurityAuditQueries lists audit queries.
ListSecurityAuditQueries(context.Context, int, string) ([]*secreports.AuditQuery, string, error)
}
// SecurityReportGetter is the interface for security report getters.
type SecurityReportGetter interface {
// GetSecurityReport returns a security report.
GetSecurityReport(ctx context.Context, name string) (*secreports.Report, error)
// GetSecurityReports returns a security report.
GetSecurityReports(ctx context.Context) ([]*secreports.Report, error)
// ListSecurityReports lists security reports.
ListSecurityReports(ctx context.Context, i int, token string) ([]*secreports.Report, string, error)
}
// SecurityReportStateGetter is the interface for security report state getters.
type SecurityReportStateGetter interface {
// GetSecurityReportState returns a security report state.
GetSecurityReportState(ctx context.Context, name string) (*secreports.ReportState, error)
// ListSecurityReportsStates lists security report states.
ListSecurityReportsStates(context.Context, int, string) ([]*secreports.ReportState, string, error)
}
// SecReports is the interface for the SecReports service.
type SecReports interface {
SecurityAuditQueryGetter
// UpsertSecurityAuditQuery upserts an audit query.
UpsertSecurityAuditQuery(ctx context.Context, in *secreports.AuditQuery) error
// DeleteSecurityAuditQuery deletes an audit query.
DeleteSecurityAuditQuery(ctx context.Context, name string) error
SecurityReportGetter
// UpsertSecurityReport upserts a security report.
UpsertSecurityReport(ctx context.Context, item *secreports.Report) error
// DeleteSecurityReport deletes a security report.
DeleteSecurityReport(ctx context.Context, name string) error
SecurityReportStateGetter
// UpsertSecurityReportsState upserts a security report state.
UpsertSecurityReportsState(ctx context.Context, item *secreports.ReportState) error
}
// CostLimiter is the interface for the security cost limiter.
type CostLimiter interface {
// UpsertCostLimiter upserts a security cost limiter.
UpsertCostLimiter(ctx context.Context, item *secreports.CostLimiter) error
// GetCostLimiter returns a security cost limiter.
GetCostLimiter(ctx context.Context, name string) (*secreports.CostLimiter, error)
}
// MarshalAuditQuery marshals an audit query.
func MarshalAuditQuery(in *secreports.AuditQuery, opts ...MarshalOption) ([]byte, error) {
if err := in.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *in
in = ©
}
return utils.FastMarshal(in)
}
// UnmarshalAuditQuery unmarshals an audit query.
func UnmarshalAuditQuery(data []byte, opts ...MarshalOption) (*secreports.AuditQuery, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing access list data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var out *secreports.AuditQuery
if err := utils.FastUnmarshal(data, &out); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := out.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.Expires.IsZero() {
out.SetExpiry(cfg.Expires)
}
return out, nil
}
// MarshalSecurityReport marshals a security report.
func MarshalSecurityReport(in *secreports.Report, opts ...MarshalOption) ([]byte, error) {
if err := in.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *in
in = ©
}
return utils.FastMarshal(in)
}
// UnmarshalSecurityReport unmarshals a security report.
func UnmarshalSecurityReport(data []byte, opts ...MarshalOption) (*secreports.Report, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing access list data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var out *secreports.Report
if err := utils.FastUnmarshal(data, &out); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := out.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.Expires.IsZero() {
out.SetExpiry(cfg.Expires)
}
return out, nil
}
// MarshalSecurityReportState marshals a security report state.
func MarshalSecurityReportState(in *secreports.ReportState, opts ...MarshalOption) ([]byte, error) {
if err := in.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *in
in = ©
}
return utils.FastMarshal(in)
}
// UnmarshalSecurityReportState unmarshals a security report state.
func UnmarshalSecurityReportState(data []byte, opts ...MarshalOption) (*secreports.ReportState, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var out *secreports.ReportState
if err := utils.FastUnmarshal(data, &out); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := out.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.Expires.IsZero() {
out.SetExpiry(cfg.Expires)
}
return out, nil
}
// MarshalSecurityCostLimiter marshals a security report state.
func MarshalSecurityCostLimiter(in *secreports.CostLimiter, opts ...MarshalOption) ([]byte, error) {
if err := in.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *in
in = ©
}
return utils.FastMarshal(in)
}
// UnmarshalSecurityCostLimiter unmarshals a security report cost limiter.
func UnmarshalSecurityCostLimiter(data []byte, opts ...MarshalOption) (*secreports.CostLimiter, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var out *secreports.CostLimiter
if err := utils.FastUnmarshal(data, &out); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := out.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.Expires.IsZero() {
out.SetExpiry(cfg.Expires)
}
return out, nil
}
/*
* 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"
"log/slog"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/retryutils"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/utils"
)
type SemaphoreLockConfig struct {
// Service is the service against which all semaphore
// operations are performed.
Service types.Semaphores
// Expiry is an optional lease expiry parameter.
Expiry time.Duration
// TickRate is the rate at which lease renewals are attempted
// and defaults to 1/2 expiry. Used to accelerate tests.
TickRate time.Duration
// Params holds the semaphore lease acquisition parameters.
Params types.AcquireSemaphoreRequest
// Clock used to alter time in tests
Clock clockwork.Clock
}
// CheckAndSetDefaults checks and sets default parameters
func (l *SemaphoreLockConfig) CheckAndSetDefaults() error {
if l.Clock == nil {
l.Clock = clockwork.NewRealClock()
}
if l.Service == nil {
return trace.BadParameter("missing semaphore service")
}
if l.Expiry == 0 {
l.Expiry = defaults.SessionControlTimeout
}
if l.Expiry < time.Millisecond {
return trace.BadParameter("sub-millisecond lease expiry is not supported: %v", l.Expiry)
}
if l.TickRate == 0 {
l.TickRate = l.Expiry / 2
}
if l.TickRate >= l.Expiry {
return trace.BadParameter("tick-rate must be less than expiry")
}
if l.Params.Expires.IsZero() {
l.Params.Expires = l.Clock.Now().UTC().Add(l.Expiry)
}
if err := l.Params.Check(); err != nil {
return trace.Wrap(err)
}
return nil
}
// SemaphoreLock provides a convenient interface for managing
// semaphore lease keepalive operations.
// SemaphoreLock implements the [context.Context] interface
// and can be used to propagate cancellation when the parent
// context is canceled or when the lease expires.
//
// lease,err := AcquireSemaphoreLock(ctx, cfg)
// if err != nil {
// ... handle error ...
// }
// defer func(){
// lease.Stop()
// err := lease.Wait()
// if err != nil {
// ... handle error ...
// }
// }()
//
// newCtx,cancel := context.WithCancel(ctx)
// defer cancel()
// ... do work with newCtx ...
type SemaphoreLock struct {
// Context is the parent context for the lease keepalive operation.
// it's used to propagate deadline cancellations from the parent
// context and to carry values for the context interface.
context.Context
cancelCtx context.CancelCauseFunc
cfg SemaphoreLockConfig
lease0 types.SemaphoreLease
retry retryutils.Retry
ticker clockwork.Ticker
closeOnce sync.Once
renewalC chan struct{}
cond *sync.Cond
err error
fin bool
}
// finish registers the final result of the background
// goroutine. must be called even if err is nil in
// order to wake any goroutines waiting on the error
// and mark the lock as finished.
func (l *SemaphoreLock) finish(err error) {
l.cond.L.Lock()
defer l.cond.L.Unlock()
l.err = err
l.fin = true
l.cond.Broadcast()
}
// Wait blocks until the final result is available. Note that
// this method may block longer than desired since cancellation of
// the parent context triggers the *start* of the release operation.
func (l *SemaphoreLock) Wait() error {
l.cond.L.Lock()
defer l.cond.L.Unlock()
for !l.fin {
l.cond.Wait()
}
return l.err
}
// Stop stops associated lease keepalive.
func (l *SemaphoreLock) Stop() {
l.closeOnce.Do(func() {
l.ticker.Stop()
l.cancelCtx(nil)
})
}
// Renewed notifies on next successful lease keepalive.
// Used in tests to block until next renewal.
func (l *SemaphoreLock) Renewed() <-chan struct{} {
return l.renewalC
}
func (l *SemaphoreLock) keepAlive() {
var nodrop bool
var err error
lease := l.lease0
defer func() {
l.cancelCtx(err)
l.Stop()
defer l.finish(err)
if nodrop {
// non-standard exit conditions; don't bother handling
// cancellation/expiry.
return
}
if lease.Expires.After(l.cfg.Clock.Now().UTC()) {
// parent context is closed. create orphan context with generous
// timeout for lease cancellation scope. this will not block any
// caller that is not explicitly waiting on the final error value.
cancelContext, cancel := context.WithTimeout(context.Background(), l.cfg.Expiry/4)
defer cancel()
err = l.cfg.Service.CancelSemaphoreLease(cancelContext, lease)
if err != nil {
slog.WarnContext(cancelContext, "Failed to cancel semaphore lease",
"semaphore_kind", lease.SemaphoreKind,
"semaphore_name", lease.SemaphoreName,
"error", err,
)
}
} else {
slog.ErrorContext(context.Background(), "Semaphore lease expired",
"semaphore_kind", lease.SemaphoreKind,
"semaphore_name", lease.SemaphoreName,
)
}
}()
Outer:
for {
select {
case tick := <-l.ticker.Chan():
leaseContext, leaseCancel := context.WithDeadline(l.Context, lease.Expires)
nextLease := lease
nextLease.Expires = tick.Add(l.cfg.Expiry)
for {
err = l.cfg.Service.KeepAliveSemaphoreLease(leaseContext, nextLease)
if trace.IsNotFound(err) {
leaseCancel()
// semaphore and/or lease no longer exist; best to log the error
// and exit immediately.
slog.WarnContext(leaseContext, "Halting keepalive on semaphore",
"semaphore_kind", lease.SemaphoreKind,
"semaphore_name", lease.SemaphoreName,
"error", err,
)
nodrop = true
return
}
if err == nil {
leaseCancel()
lease = nextLease
l.retry.Reset()
select {
case l.renewalC <- struct{}{}:
default:
}
continue Outer
}
slog.DebugContext(leaseContext, "Failed to renew semaphore lease",
"semaphore_kind", lease.SemaphoreKind,
"semaphore_name", lease.SemaphoreName,
"error", err,
)
l.retry.Inc()
select {
case <-l.retry.After():
case tick = <-l.ticker.Chan():
// check to make sure that we still have some time on the lease. the default tick rate would have
// us waking _as_ the lease expires here, but if we're working with a higher tick rate, its worth
// retrying again.
if !lease.Expires.After(tick) {
leaseCancel()
return
}
case <-leaseContext.Done():
leaseCancel() // demanded by linter
return
case <-l.Done():
leaseCancel()
return
}
}
case <-l.Done():
return
}
}
}
// AcquireSemaphoreWithRetryConfig contains parameters for trying to acquire a
// semaphore with a retry.
type AcquireSemaphoreWithRetryConfig struct {
Service types.Semaphores
Request types.AcquireSemaphoreRequest
Retry retryutils.LinearConfig
// TTL, if set, will be used to set the expiry of the request.
TTL time.Duration
// Now, if set, will be used instead of time.Now when calculating the expiry
// of the request.
Now func() time.Time
}
// AcquireSemaphoreWithRetry tries to acquire the semaphore according to the
// retry schedule until it succeeds or context expires.
func AcquireSemaphoreWithRetry(ctx context.Context, req AcquireSemaphoreWithRetryConfig) (*types.SemaphoreLease, error) {
retry, err := retryutils.NewLinear(req.Retry)
if err != nil {
return nil, trace.Wrap(err)
}
if req.Now == nil {
req.Now = time.Now
}
var lease *types.SemaphoreLease
err = retry.For(ctx, func() (err error) {
r := req.Request
if req.TTL > 0 {
r.Expires = req.Now().Add(req.TTL)
}
lease, err = req.Service.AcquireSemaphore(ctx, r)
return trace.Wrap(err)
})
if err != nil {
return nil, trace.Wrap(err)
}
return lease, nil
}
// AcquireSemaphoreLock attempts to acquire and hold a semaphore lease. If successfully acquired,
// background keepalive processes are started and an associated lock handle is returned. Canceling
// the supplied context releases the semaphore.
func AcquireSemaphoreLock(ctx context.Context, cfg SemaphoreLockConfig) (*SemaphoreLock, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
// set up retry with a ratio which will result in 3-4 retries before the lease expires
retry, err := retryutils.NewLinear(retryutils.LinearConfig{
Max: cfg.Expiry / 4,
Step: cfg.Expiry / 16,
Jitter: retryutils.DefaultJitter,
Clock: cfg.Clock,
})
if err != nil {
return nil, trace.Wrap(err)
}
lease, err := cfg.Service.AcquireSemaphore(ctx, cfg.Params)
if err != nil {
return nil, trace.Wrap(err)
}
ctx, cancel := context.WithCancelCause(ctx)
lock := &SemaphoreLock{
Context: ctx,
cancelCtx: cancel,
cfg: cfg,
lease0: *lease,
retry: retry,
ticker: cfg.Clock.NewTicker(cfg.TickRate),
renewalC: make(chan struct{}),
cond: sync.NewCond(&sync.Mutex{}),
}
go lock.keepAlive()
return lock, nil
}
// SemaphoreLockConfigWithRetry contains parameters for acquiring a semaphore lock
// until it succeeds or context expires.
type SemaphoreLockConfigWithRetry struct {
SemaphoreLockConfig
// Retry is the retry configuration.
Retry retryutils.LinearConfig
}
// AcquireSemaphoreLockWithRetry attempts to acquire and hold a semaphore lease. If successfully acquired,
// background keepalive processes are started and an associated lock handle is returned.
// If the lease cannot be acquired, the operation is retried according to the retry schedule until
// it succeeds or the context expires. Canceling the supplied context releases the semaphore.
func AcquireSemaphoreLockWithRetry(ctx context.Context, cfg SemaphoreLockConfigWithRetry) (*SemaphoreLock, error) {
retry, err := retryutils.NewLinear(cfg.Retry)
if err != nil {
return nil, trace.Wrap(err)
}
var lease *SemaphoreLock
err = retry.For(ctx, func() (err error) {
lease, err = AcquireSemaphoreLock(ctx, cfg.SemaphoreLockConfig)
return trace.Wrap(err)
})
if err != nil {
return nil, trace.Wrap(err)
}
return lease, nil
}
// UnmarshalSemaphore unmarshals the Semaphore resource from JSON.
func UnmarshalSemaphore(bytes []byte, opts ...MarshalOption) (types.Semaphore, error) {
var semaphore types.SemaphoreV3
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &semaphore); err != nil {
return nil, trace.BadParameter("%s", err)
}
err = semaphore.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
semaphore.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
semaphore.SetExpiry(cfg.Expires)
}
return &semaphore, nil
}
// MarshalSemaphore marshals the Semaphore resource to JSON.
func MarshalSemaphore(semaphore types.Semaphore, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch semaphore := semaphore.(type) {
case *types.SemaphoreV3:
if err := semaphore.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, semaphore))
default:
return nil, trace.BadParameter("unrecognized resource version %T", semaphore)
}
}
/*
* 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 (
"encoding/json"
"fmt"
"maps"
"slices"
"time"
"github.com/gravitational/trace"
apidefaults "github.com/gravitational/teleport/api/defaults"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/wrappers"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils"
)
const (
// Equal means two objects are equal
Equal = iota
// OnlyTimestampsDifferent is true when only timestamps are different
OnlyTimestampsDifferent = iota
// Different means that some fields are different
Different = iota
)
// CompareServers compares two provided servers.
func CompareServers(a, b types.Resource) int {
if serverA, ok := a.(types.Server); ok {
if serverB, ok := b.(types.Server); ok {
return compareServers(serverA, serverB)
}
}
if appA, ok := a.(types.AppServer); ok {
if appB, ok := b.(types.AppServer); ok {
return compareApplicationServers(appA, appB)
}
}
if kubeA, ok := a.(types.KubeServer); ok {
if kubeB, ok := b.(types.KubeServer); ok {
return compareKubernetesServers(kubeA, kubeB)
}
}
if dbA, ok := a.(types.DatabaseServer); ok {
if dbB, ok := b.(types.DatabaseServer); ok {
return compareDatabaseServers(dbA, dbB)
}
}
if dbServiceA, ok := a.(types.DatabaseService); ok {
if dbServiceB, ok := b.(types.DatabaseService); ok {
return compareDatabaseServices(dbServiceA, dbServiceB)
}
}
if winA, ok := a.(types.WindowsDesktopService); ok {
if winB, ok := b.(types.WindowsDesktopService); ok {
return compareWindowsDesktopServices(winA, winB)
}
}
if linA, ok := a.(types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop]); ok {
if linB, ok := b.(types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop]); ok {
return CompareLinuxDesktop(linA.UnwrapT(), linB.UnwrapT())
}
}
return Different
}
func compareServers(a, b types.Server) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetName() != b.GetName() {
return Different
}
if a.GetAddr() != b.GetAddr() {
return Different
}
if a.GetHostname() != b.GetHostname() {
return Different
}
if a.GetNamespace() != b.GetNamespace() {
return Different
}
if len(a.GetPublicAddrs()) != len(b.GetPublicAddrs()) {
return Different
}
if !slices.Equal(a.GetPublicAddrs(), b.GetPublicAddrs()) {
return Different
}
r := a.GetRotation()
if !r.Matches(b.GetRotation()) {
return Different
}
if a.GetUseTunnel() != b.GetUseTunnel() {
return Different
}
if !maps.Equal(a.GetStaticLabels(), b.GetStaticLabels()) {
return Different
}
if !maps.EqualFunc(a.GetCmdLabels(), b.GetCmdLabels(), func(label types.CommandLabel, label2 types.CommandLabel) bool {
return slices.Equal(label.GetCommand(), label2.GetCommand()) &&
label.GetPeriod() == label2.GetPeriod() &&
label.GetResult() == label2.GetResult()
}) {
return Different
}
if !maps.Equal(a.GetImmutableLabels(), b.GetImmutableLabels()) {
return Different
}
if a.GetTeleportVersion() != b.GetTeleportVersion() {
return Different
}
if !slices.Equal(a.GetProxyIDs(), b.GetProxyIDs()) {
return Different
}
if a.GetRelayGroup() != b.GetRelayGroup() {
return Different
}
if !slices.Equal(a.GetRelayIDs(), b.GetRelayIDs()) {
return Different
}
if (a.GetGitHub() == nil && b.GetGitHub() != nil) ||
(a.GetGitHub() != nil && b.GetGitHub() == nil) {
return Different
}
if a.GetGitHub() != nil && b.GetGitHub() != nil {
if a.GetGitHub().Integration != b.GetGitHub().Integration {
return Different
}
if a.GetGitHub().Organization != b.GetGitHub().Organization {
return Different
}
}
if a.GetScope() != b.GetScope() {
return Different
}
if !maps.Equal(a.GetAllLabels(), b.GetAllLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.Expiry().Equal(b.Expiry()) {
return OnlyTimestampsDifferent
}
return Equal
}
func compareApplicationServers(a, b types.AppServer) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetName() != b.GetName() {
return Different
}
if a.GetNamespace() != b.GetNamespace() {
return Different
}
if a.GetTeleportVersion() != b.GetTeleportVersion() {
return Different
}
r := a.GetRotation()
if !r.Matches(b.GetRotation()) {
return Different
}
if !a.GetApp().IsEqual(b.GetApp()) {
return Different
}
if !slices.Equal(a.GetProxyIDs(), b.GetProxyIDs()) {
return Different
}
if a.GetRelayGroup() != b.GetRelayGroup() {
return Different
}
if !slices.Equal(a.GetRelayIDs(), b.GetRelayIDs()) {
return Different
}
if a.GetScope() != b.GetScope() {
return Different
}
if !maps.Equal(a.GetAllLabels(), b.GetAllLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.Expiry().Equal(b.Expiry()) {
return OnlyTimestampsDifferent
}
return Equal
}
func compareDatabaseServices(a, b types.DatabaseService) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetName() != b.GetName() {
return Different
}
if a.GetNamespace() != b.GetNamespace() {
return Different
}
if !slices.EqualFunc(a.GetResourceMatchers(), b.GetResourceMatchers(),
func(matcher *types.DatabaseResourceMatcher, matcher2 *types.DatabaseResourceMatcher) bool {
return matcher.AWS.AssumeRoleARN == matcher2.AWS.AssumeRoleARN &&
maps.EqualFunc(matcher.Labels.ToProto().Values, matcher2.Labels.ToProto().Values,
func(values wrappers.StringValues, values2 wrappers.StringValues) bool {
return slices.Equal(values.Values, values2.Values)
})
}) {
return Different
}
if !maps.Equal(a.GetAllLabels(), b.GetAllLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.Expiry().Equal(b.Expiry()) {
return OnlyTimestampsDifferent
}
return Equal
}
func compareKubernetesServers(a, b types.KubeServer) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetName() != b.GetName() {
return Different
}
if a.GetNamespace() != b.GetNamespace() {
return Different
}
if a.GetTeleportVersion() != b.GetTeleportVersion() {
return Different
}
r := a.GetRotation()
if !r.Matches(b.GetRotation()) {
return Different
}
if !a.GetCluster().IsEqual(b.GetCluster()) {
return Different
}
if !slices.Equal(a.GetProxyIDs(), b.GetProxyIDs()) {
return Different
}
if a.GetRelayGroup() != b.GetRelayGroup() {
return Different
}
if !slices.Equal(a.GetRelayIDs(), b.GetRelayIDs()) {
return Different
}
if a.GetScope() != b.GetScope() {
return Different
}
if !maps.Equal(a.GetAllLabels(), b.GetAllLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.Expiry().Equal(b.Expiry()) {
return OnlyTimestampsDifferent
}
return Equal
}
func compareDatabaseServers(a, b types.DatabaseServer) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetName() != b.GetName() {
return Different
}
if a.GetNamespace() != b.GetNamespace() {
return Different
}
if a.GetTeleportVersion() != b.GetTeleportVersion() {
return Different
}
r := a.GetRotation()
if !r.Matches(b.GetRotation()) {
return Different
}
if !a.GetDatabase().IsEqual(b.GetDatabase()) {
return Different
}
if !slices.Equal(a.GetProxyIDs(), b.GetProxyIDs()) {
return Different
}
if a.GetRelayGroup() != b.GetRelayGroup() {
return Different
}
if !slices.Equal(a.GetRelayIDs(), b.GetRelayIDs()) {
return Different
}
if a.GetScope() != b.GetScope() {
return Different
}
if !maps.Equal(a.GetAllLabels(), b.GetAllLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.Expiry().Equal(b.Expiry()) {
return OnlyTimestampsDifferent
}
return Equal
}
func compareWindowsDesktopServices(a, b types.WindowsDesktopService) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetName() != b.GetName() {
return Different
}
if a.GetAddr() != b.GetAddr() {
return Different
}
if a.GetTeleportVersion() != b.GetTeleportVersion() {
return Different
}
if !slices.Equal(a.GetProxyIDs(), b.GetProxyIDs()) {
return Different
}
if a.GetRelayGroup() != b.GetRelayGroup() {
return Different
}
if !slices.Equal(a.GetRelayIDs(), b.GetRelayIDs()) {
return Different
}
if !maps.Equal(a.GetAllLabels(), b.GetAllLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.Expiry().Equal(b.Expiry()) {
return OnlyTimestampsDifferent
}
return Equal
}
func CompareLinuxDesktop(a, b *linuxdesktopv1.LinuxDesktop) int {
if a.GetKind() != b.GetKind() {
return Different
}
if a.GetSubKind() != b.GetSubKind() {
return Different
}
if a.GetMetadata().GetName() != b.GetMetadata().GetName() {
return Different
}
if a.GetSpec().GetAddr() != b.GetSpec().GetAddr() {
return Different
}
if a.GetSpec().GetHostname() != b.GetSpec().GetHostname() {
return Different
}
if !slices.Equal(a.GetSpec().GetProxyIds(), b.GetSpec().GetProxyIds()) {
return Different
}
if !maps.Equal(a.GetMetadata().GetLabels(), b.GetMetadata().GetLabels()) {
return Different
}
// OnlyTimestampsDifferent check must be after all Different checks.
if !a.GetMetadata().GetExpires().AsTime().Equal(b.GetMetadata().GetExpires().AsTime()) {
return OnlyTimestampsDifferent
}
return Equal
}
// CommandLabels is a set of command labels
type CommandLabels map[string]types.CommandLabel
// Clone returns copy of the set
func (c *CommandLabels) Clone() CommandLabels {
out := make(CommandLabels, len(*c))
for name, label := range *c {
out[name] = label.Clone()
}
return out
}
// SetEnv sets the value of the label from environment variable
func (c *CommandLabels) SetEnv(v string) error {
if err := json.Unmarshal([]byte(v), c); err != nil {
return trace.Wrap(err, "can not parse Command Labels")
}
return nil
}
// SortedServers is a sort wrapper that sorts servers by name
type SortedServers []types.Server
func (s SortedServers) Len() int {
return len(s)
}
func (s SortedServers) Less(i, j int) bool {
return s[i].GetName() < s[j].GetName()
}
func (s SortedServers) Swap(i, j int) {
s[i], s[j] = s[j], s[i]
}
// SortedReverseTunnels sorts reverse tunnels by cluster name
type SortedReverseTunnels []types.ReverseTunnel
func (s SortedReverseTunnels) Len() int {
return len(s)
}
func (s SortedReverseTunnels) Less(i, j int) bool {
return s[i].GetClusterName() < s[j].GetClusterName()
}
func (s SortedReverseTunnels) Swap(i, j int) {
s[i], s[j] = s[j], s[i]
}
// GuessProxyHostAndVersion tries to find the first proxy with a public
// address configured and return that public addr and version.
// If no proxies are configured, it will return a guessed value by concatenating
// the first proxy's hostname with default port number, and the first proxy's
// version will also be returned.
//
// Returns empty value if there are no proxies.
func GuessProxyHostAndVersion(proxies []types.Server) (string, string, error) {
if len(proxies) == 0 {
return "", "", trace.NotFound("list of proxies empty")
}
// Find the first proxy with a public address set and return it.
for _, proxy := range proxies {
proxyHost := proxy.GetPublicAddr()
if proxyHost != "" {
return proxyHost, proxy.GetTeleportVersion(), nil
}
}
// No proxies have a public address set, return guessed value.
guessProxyHost := fmt.Sprintf("%v:%v", proxies[0].GetHostname(), defaults.HTTPListenPort)
return guessProxyHost, proxies[0].GetTeleportVersion(), nil
}
// UnmarshalServer unmarshals the Server resource from JSON.
func UnmarshalServer(bytes []byte, kind string, opts ...MarshalOption) (types.Server, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if len(bytes) == 0 {
return nil, trace.BadParameter("missing server data")
}
var s types.ServerV2
if err := utils.FastUnmarshal(bytes, &s); err != nil {
return nil, trace.BadParameter("%s", err)
}
s.Kind = kind
if err := s.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
s.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
s.SetExpiry(cfg.Expires)
}
if s.Metadata.Expires != nil {
apiutils.UTC(s.Metadata.Expires)
}
// Force the timestamps to UTC for consistency.
// See https://github.com/gogo/protobuf/issues/519 for details on issues this causes for proto.Clone
apiutils.UTC(&s.Spec.Rotation.Started)
apiutils.UTC(&s.Spec.Rotation.LastRotated)
return &s, nil
}
// MarshalServer marshals the Server resource to JSON.
func MarshalServer(server types.Server, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch server := server.(type) {
case *types.ServerV2:
if err := server.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, server))
default:
return nil, trace.BadParameter("unrecognized server version %T", server)
}
}
// UnmarshalServers unmarshals a list of Server resources.
func UnmarshalServers(bytes []byte) ([]types.Server, error) {
var servers []types.ServerV2
err := utils.FastUnmarshal(bytes, &servers)
if err != nil {
return nil, trace.Wrap(err)
}
out := make([]types.Server, len(servers))
for i := range servers {
out[i] = types.Server(&servers[i])
}
return out, nil
}
// MarshalServers marshals a list of Server resources.
func MarshalServers(s []types.Server) ([]byte, error) {
bytes, err := utils.FastMarshal(s)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
}
// GetCursorForNode returns the resource cursor identifying a node in the
// logical resource stream: "<name>" for unscoped nodes and
// "~scoped/<encoded-scope>/<name>" for scoped nodes.
func GetCursorForNode(server types.Server) string {
return scopes.MakeResourceCursor(server.GetScope(), server.GetName())
}
// NodeHasMissedKeepAlives checks if node has missed its keep alive
func NodeHasMissedKeepAlives(s types.Server) bool {
serverExpiry := s.Expiry()
return serverExpiry.Before(time.Now().Add(apidefaults.ServerAnnounceTTL - (apidefaults.ServerKeepAliveTTL() * 2)))
}
// EqualFromBool is a helper function that converts a boolean value to an integer
// value that represents the equality status.
func EqualFromBool(b bool) int {
if !b {
return Different
}
return Equal
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalServerInfo unmarshals the ServerInfo resource from JSON.
func UnmarshalServerInfo(bytes []byte, opts ...MarshalOption) (types.ServerInfo, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if len(bytes) == 0 {
return nil, trace.BadParameter("missing server info data")
}
var si types.ServerInfoV1
if err := utils.FastUnmarshal(bytes, &si); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := si.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
si.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
si.SetExpiry(cfg.Expires)
}
if si.Metadata.Expires != nil {
apiutils.UTC(si.Metadata.Expires)
}
return &si, nil
}
// MarshalServerInfo marshals the ServerInfo resource to JSON.
func MarshalServerInfo(si types.ServerInfo, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch si := si.(type) {
case *types.ServerInfoV1:
if err := si.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
bytes, err := utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, si))
return bytes, trace.Wrap(err)
default:
return nil, trace.BadParameter("unrecognized server info version %T", si)
}
}
// UnmarshalServerInfos unmarshals a list of ServerInfo resources.
func UnmarshalServerInfos(bytes []byte) ([]types.ServerInfo, error) {
var serverInfos []types.ServerInfoV1
err := utils.FastUnmarshal(bytes, &serverInfos)
if err != nil {
return nil, trace.Wrap(err)
}
out := make([]types.ServerInfo, len(serverInfos))
for i := range serverInfos {
out[i] = types.ServerInfo(&serverInfos[i])
}
return out, nil
}
// MarshalServerInfos marshals a list of ServerInfo resources.
func MarshalServerInfos(si []types.ServerInfo) ([]byte, error) {
bytes, err := utils.FastMarshal(si)
if err != nil {
return nil, trace.Wrap(err)
}
return bytes, nil
}
/*
* 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"
"encoding/json"
"iter"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalWebSession unmarshals the WebSession resource from JSON.
func UnmarshalWebSession(bytes []byte, opts ...MarshalOption) (types.WebSession, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
err = json.Unmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V2:
var ws types.WebSessionV2
if err := utils.FastUnmarshal(bytes, &ws); err != nil {
return nil, trace.Wrap(err)
}
apiutils.UTC(&ws.Spec.BearerTokenExpires)
apiutils.UTC(&ws.Spec.Expires)
if err := ws.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
ws.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
ws.SetExpiry(cfg.Expires)
}
return &ws, nil
}
return nil, trace.BadParameter("web session resource version %v is not supported", h.Version)
}
// MarshalWebSession marshals the WebSession resource to JSON.
func MarshalWebSession(webSession types.WebSession, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch webSession := webSession.(type) {
case *types.WebSessionV2:
if version := webSession.GetVersion(); version != types.V2 {
return nil, trace.BadParameter("mismatched web session version %v and type %T", version, webSession)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, webSession))
default:
return nil, trace.BadParameter("unrecognized web session version %T", webSession)
}
}
// WebToken defines an interface for managing web token resources.
type WebToken interface {
// GetWebToken gets a web token.
GetWebToken(context.Context, types.GetWebTokenRequest) (types.WebToken, error)
// GetWebTokens gets all web tokens.
GetWebTokens(context.Context) ([]types.WebToken, error)
// ListWebTokens returns a page of web tokens
ListWebTokens(ctx context.Context, limit int, start string) ([]types.WebToken, string, error)
// RangeWebTokens returns web tokens within the range [start, end).
RangeWebTokens(ctx context.Context, start, end string) iter.Seq2[types.WebToken, error]
// UpsertWebToken updates the existing or inserts a new web token.
UpsertWebToken(context.Context, types.WebToken) error
// DeleteWebToken deletes a web token.
DeleteWebToken(context.Context, types.DeleteWebTokenRequest) error
// DeleteAllWebTokens deletes all web tokens.
DeleteAllWebTokens(context.Context) error
}
// MarshalWebToken serializes the web token as JSON-encoded payload
func MarshalWebToken(webToken types.WebToken, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch webToken := webToken.(type) {
case *types.WebTokenV3:
if err := webToken.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, webToken))
default:
return nil, trace.BadParameter("unrecognized web token version %T", webToken)
}
}
// UnmarshalWebToken interprets bytes as JSON-encoded web token value
func UnmarshalWebToken(bytes []byte, opts ...MarshalOption) (types.WebToken, error) {
config, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var hdr types.ResourceHeader
err = json.Unmarshal(bytes, &hdr)
if err != nil {
return nil, trace.Wrap(err)
}
switch hdr.Version {
case types.V3:
var token types.WebTokenV3
if err := utils.FastUnmarshal(bytes, &token); err != nil {
return nil, trace.BadParameter("invalid web token: %v", err.Error())
}
if err := token.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if config.Revision != "" {
token.SetRevision(config.Revision)
}
if !config.Expires.IsZero() {
token.Metadata.SetExpiry(config.Expires)
}
apiutils.UTC(token.Metadata.Expires)
return &token, nil
}
return nil, trace.BadParameter("web token resource version %v is not supported", hdr.Version)
}
/*
* Teleport
* Copyright (C) 2025 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 (
apidefaults "github.com/gravitational/teleport/api/defaults"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
"github.com/gravitational/teleport/api/types"
apievents "github.com/gravitational/teleport/api/types/events"
)
// ExtendWithSessionEnd extends the context with a session end event and
// rebuilds the resource from the event. An AccessChecker must be provided
// to allow access checks to other resources in the where clause.
func (ctx *Context) ExtendWithSessionEnd(sessionEnd apievents.AuditEvent, checker AccessChecker) {
ctx.Session = sessionEnd
ctx.Resource = rebuildResourceFromSessionEndEvent(sessionEnd)
if linuxEnd, ok := sessionEnd.(*apievents.LinuxDesktopSessionEnd); ok {
ctx.Resource153 = linuxdesktopv1.LinuxDesktop_builder{
Kind: types.KindLinuxDesktop,
SubKind: "",
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: linuxEnd.DesktopName,
Labels: linuxEnd.DesktopLabels,
}.Build(),
Spec: linuxdesktopv1.LinuxDesktopSpec_builder{
Addr: linuxEnd.DesktopAddr,
}.Build(),
}.Build()
}
// AccessCheker is set here to allow access checks to other resources
// in the where clause.
ctx.AccessChecker = checker
}
// rebuildResourceFromSessionEndEvent rebuilds a resource from a session end event.
// This is used to reconstruct the resource that was active at the time of the session end event
// for audit log RBAC purposes.
func rebuildResourceFromSessionEndEvent(event apievents.AuditEvent) types.Resource {
switch sEnd := event.(type) {
case *apievents.SessionEnd:
if sEnd == nil {
return nil
}
switch sEnd.Protocol {
case apievents.EventProtocolSSH:
return &types.ServerV2{
Kind: types.KindNode,
Version: types.V2,
Metadata: types.Metadata{
Name: sEnd.ServerMetadata.ServerID,
Namespace: sEnd.ServerMetadata.ServerNamespace,
Labels: sEnd.ServerMetadata.ServerLabels,
},
Spec: types.ServerSpecV2{
Addr: sEnd.ServerMetadata.ServerAddr,
Hostname: sEnd.ServerMetadata.ServerHostname,
},
}
case apievents.EventProtocolKube:
return &types.KubernetesClusterV3{
Kind: types.KindKubernetesCluster,
Version: types.V3,
Metadata: types.Metadata{
Name: sEnd.KubernetesClusterMetadata.KubernetesCluster,
Namespace: apidefaults.Namespace,
Labels: sEnd.KubernetesClusterMetadata.KubernetesLabels,
},
Spec: types.KubernetesClusterSpecV3{},
}
}
case *apievents.WindowsDesktopSessionEnd:
if sEnd == nil {
return nil
}
return &types.WindowsDesktopV3{
ResourceHeader: types.ResourceHeader{
Kind: types.KindWindowsDesktop,
Version: types.V3,
Metadata: types.Metadata{
Name: sEnd.DesktopName,
Namespace: apidefaults.Namespace,
Labels: sEnd.DesktopLabels,
},
},
Spec: types.WindowsDesktopSpecV3{
Addr: sEnd.DesktopAddr,
Domain: sEnd.Domain,
},
}
case *apievents.AppSessionChunk:
if sEnd == nil {
return nil
}
return &types.AppV3{
Kind: types.KindApp,
Version: types.V3,
Metadata: types.Metadata{
Name: sEnd.AppName,
Namespace: apidefaults.Namespace,
Labels: sEnd.AppLabels,
},
Spec: types.AppSpecV3{
URI: sEnd.AppURI,
PublicAddr: sEnd.AppPublicAddr,
},
}
case *apievents.DatabaseSessionEnd:
if sEnd == nil {
return nil
}
return &types.DatabaseV3{
Kind: types.KindDatabase,
Version: types.V3,
Metadata: types.Metadata{
Name: sEnd.DatabaseService,
Namespace: apidefaults.Namespace,
Labels: sEnd.DatabaseLabels,
},
Spec: types.DatabaseSpecV3{
Protocol: sEnd.DatabaseProtocol,
URI: sEnd.DatabaseURI,
},
}
}
return nil
}
/*
* 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 (
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// IsRecordAtProxy returns true if recording is sync or async at proxy.
func IsRecordAtProxy(mode string) bool {
return mode == types.RecordAtProxy || mode == types.RecordAtProxySync
}
// IsRecordSync returns true if recording is sync for proxy or node.
func IsRecordSync(mode string) bool {
return mode == types.RecordAtProxySync || mode == types.RecordAtNodeSync
}
// UnmarshalSessionRecordingConfig unmarshals the SessionRecordingConfig resource from JSON.
func UnmarshalSessionRecordingConfig(bytes []byte, opts ...MarshalOption) (types.SessionRecordingConfig, error) {
var recConfig types.SessionRecordingConfigV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &recConfig); err != nil {
return nil, trace.BadParameter("%s", err)
}
err = recConfig.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
recConfig.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
recConfig.SetExpiry(cfg.Expires)
}
return &recConfig, nil
}
// MarshalSessionRecordingConfig marshals the SessionRecordingConfig resource to JSON.
func MarshalSessionRecordingConfig(recConfig types.SessionRecordingConfig, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch recConfig := recConfig.(type) {
case *types.SessionRecordingConfigV2:
if err := recConfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if version := recConfig.GetVersion(); version != types.V2 {
return nil, trace.BadParameter("mismatched session recording config version %v and type %T", version, recConfig)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, recConfig))
default:
return nil, trace.BadParameter("unrecognized session recording config version %T", recConfig)
}
}
const (
KeyTypeAWS = "aws_kms"
KeyTypeGCP = "gcp_kms"
KeyTypePKCS11 = "pkcs11"
KeyTypeSoftware = "software"
)
var errRecordingEncryptionWithFIPS = &trace.BadParameterError{Message: `non-FIPS compliant session recording setting: "encryption" must not be enabled`}
var errManualKeyManagementInCloud = &trace.BadParameterError{Message: `"manual_key_management" configuration is unsupported in Teleport Cloud`}
// ValidateSessionRecordingConfig checks that the state of a [SessionRecordingConfig] meets constraints.
func ValidateSessionRecordingConfig(cfg types.SessionRecordingConfig, fips, cloud bool) error {
if !slices.Contains(types.SessionRecordingModes, cfg.GetMode()) {
return trace.BadParameter("session recording mode must be one of %v; got %q", strings.Join(types.SessionRecordingModes, ","), cfg.GetMode())
}
encryptionCfg := cfg.GetEncryptionConfig()
if encryptionCfg == nil || !encryptionCfg.Enabled {
return nil
}
if fips && encryptionCfg.Enabled {
return trace.Wrap(errRecordingEncryptionWithFIPS)
}
manualKeyManagement := encryptionCfg.ManualKeyManagement
if manualKeyManagement == nil || !manualKeyManagement.Enabled {
return nil
}
if cloud {
return trace.Wrap(errManualKeyManagementInCloud)
}
if len(manualKeyManagement.ActiveKeys) == 0 {
return trace.BadParameter("at least one active key must be configured when using manually managed encryption keys")
}
for _, label := range manualKeyManagement.ActiveKeys {
switch strings.ToLower(label.Type) {
case KeyTypeAWS, KeyTypeGCP, KeyTypePKCS11, KeyTypeSoftware:
default:
return trace.BadParameter("invalid key type %q found for active manually managed key", label.Type)
}
}
for _, label := range manualKeyManagement.RotatedKeys {
switch strings.ToLower(label.Type) {
case KeyTypeAWS, KeyTypeGCP, KeyTypePKCS11, KeyTypeSoftware:
default:
return trace.BadParameter("invalid key type %q found for rotated manually managed key", label.Type)
}
}
return nil
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// SessionTrackerService is a realtime session service that has information about
// sessions that are in-flight in the cluster at the moment.
type SessionTrackerService interface {
// GetActiveSessionTrackers returns a list of active session trackers.
GetActiveSessionTrackers(ctx context.Context) ([]types.SessionTracker, error)
// GetActiveSessionTrackersWithFilter returns a list of active sessions filtered by a filter.
GetActiveSessionTrackersWithFilter(ctx context.Context, filter *types.SessionTrackerFilter) ([]types.SessionTracker, error)
// GetSessionTracker returns the current state of a session tracker for an active session.
GetSessionTracker(ctx context.Context, sessionID string) (types.SessionTracker, error)
// CreateSessionTracker creates a tracker resource for an active session.
CreateSessionTracker(ctx context.Context, st types.SessionTracker) (types.SessionTracker, error)
// UpdateSessionTracker updates a tracker resource for an active session.
UpdateSessionTracker(ctx context.Context, req *proto.UpdateSessionTrackerRequest) error
// RemoveSessionTracker removes a tracker resource for an active session.
RemoveSessionTracker(ctx context.Context, sessionID string) error
// UpdatePresence updates the presence status of a user in a session.
UpdatePresence(ctx context.Context, sessionID, user, userCluster string) error
}
// UnmarshalSessionTracker unmarshals the Session resource from JSON.
func UnmarshalSessionTracker(bytes []byte) (types.SessionTracker, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
var session types.SessionTrackerV1
if err := utils.FastUnmarshal(bytes, &session); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := session.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &session, nil
}
// MarshalSessionTracker marshals the Session resource to JSON.
func MarshalSessionTracker(session types.SessionTracker) ([]byte, error) {
switch session := session.(type) {
case *types.SessionTrackerV1:
if err := session.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(session)
default:
return nil, trace.BadParameter("unrecognized session version %T", session)
}
}
// Teleport
// Copyright (C) 2025 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 (
"bytes"
"context"
"crypto/x509"
"encoding/pem"
"regexp"
"strings"
"github.com/gravitational/trace"
"github.com/sigstore/sigstore-go/pkg/root"
workloadidentityv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/api/types"
)
// SigstorePolicies is an interface over the SigstorePolicy service. This
// interface may also be implemented by a client to allow remote and local
// consumers to access the resource in a similar way.
type SigstorePolicies interface {
// GetSigstorePolicy gets a SigstorePolicy by name.
GetSigstorePolicy(
ctx context.Context, name string,
) (*workloadidentityv1pb.SigstorePolicy, error)
// ListSigtorePolicies lists all SigstorePolicy resources using Google style
// pagination.
ListSigstorePolicies(
ctx context.Context, pageSize int, lastToken string,
) ([]*workloadidentityv1pb.SigstorePolicy, string, error)
// CreateSigstorePolicy creates a new SigstorePolicy.
CreateSigstorePolicy(
ctx context.Context,
sigstorePolicy *workloadidentityv1pb.SigstorePolicy,
) (*workloadidentityv1pb.SigstorePolicy, error)
// DeleteSigstorePolicy deletes a SigstorePolicy by name.
DeleteSigstorePolicy(ctx context.Context, name string) error
// UpdateSigstorePolicy updates a specific SigstorePolicy. The resource must
// already exist, and, conditional update semantics are used - e.g the
// submitted resource must have a revision matching the revision of the
// resource in the backend.
UpdateSigstorePolicy(
ctx context.Context,
sigstorePolicy *workloadidentityv1pb.SigstorePolicy,
) (*workloadidentityv1pb.SigstorePolicy, error)
// UpsertSigstorePolicy creates or updates a SigstorePolicy.
UpsertSigstorePolicy(
ctx context.Context,
sigstorePolicy *workloadidentityv1pb.SigstorePolicy,
) (*workloadidentityv1pb.SigstorePolicy, error)
}
// MarshalSigstorePolicy marshals the SigstorePolicy object into a JSON byte
// slice.
func MarshalSigstorePolicy(
object *workloadidentityv1pb.SigstorePolicy, opts ...MarshalOption,
) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalSigstorePolicy unmarshals the SigstorePolicy object from a JSON byte
// slice.
func UnmarshalSigstorePolicy(
data []byte, opts ...MarshalOption,
) (*workloadidentityv1pb.SigstorePolicy, error) {
return UnmarshalProtoResource[*workloadidentityv1pb.SigstorePolicy](data, opts...)
}
// ValidateSigstorePolicy validates the SigstorePolicy object.
func ValidateSigstorePolicy(s *workloadidentityv1pb.SigstorePolicy) error {
switch {
case s.GetKind() != types.KindSigstorePolicy:
return trace.BadParameter("kind: must be %q", types.KindSigstorePolicy)
case s.GetSubKind() != "":
return trace.BadParameter("sub_kind: must be empty")
case s.GetVersion() != types.V1:
return trace.BadParameter("version: only %q is supported", types.V1)
case s.GetMetadata() == nil:
return trace.BadParameter("metadata: is required")
case s.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name: is required")
case s.GetSpec() == nil:
return trace.BadParameter("spec: is required")
case s.GetSpec().GetKey() == nil && s.GetSpec().GetKeyless() == nil:
return trace.BadParameter("spec.authority: key or keyless authority is required")
case s.GetSpec().GetRequirements() == nil:
return trace.BadParameter("spec.requirements: is required")
}
switch s.GetSpec().WhichAuthority() {
case workloadidentityv1pb.SigstorePolicySpec_Key_case:
public := s.GetSpec().GetKey().GetPublic()
if public == "" {
return trace.BadParameter("spec.key.public: is required")
}
block, rest := pem.Decode([]byte(public))
if block == nil {
return trace.BadParameter("spec.key.public: is not PEM encoded")
}
if !strings.Contains(block.Type, "PUBLIC KEY") {
return trace.BadParameter("spec.key.public: must contain a public key, not: '%s'", block.Type)
}
if len(bytes.TrimSpace(rest)) != 0 {
return trace.BadParameter("spec.key.public: must contain exactly one public key")
}
if _, err := x509.ParsePKIXPublicKey(block.Bytes); err != nil {
return trace.BadParameter("spec.key.public: failed to parse public key: %v", err)
}
case workloadidentityv1pb.SigstorePolicySpec_Keyless_case:
if len(s.GetSpec().GetKeyless().GetIdentities()) == 0 {
return trace.BadParameter("spec.keyless.identities: at least one trusted identity is required")
}
for idx, identity := range s.GetSpec().GetKeyless().GetIdentities() {
switch {
case identity.GetIssuer() != "":
case identity.GetIssuerRegex() != "":
if _, err := regexp.Compile(identity.GetIssuerRegex()); err != nil {
return trace.BadParameter("spec.keyless.identities[%d].issuer_regex: failed to parse regex: %v", idx, err)
}
default:
return trace.BadParameter("spec.keyless.identities[%d].issuer_matcher: issuer or issuer_regex is required", idx)
}
switch {
case identity.GetSubject() != "":
case identity.GetSubjectRegex() != "":
if _, err := regexp.Compile(identity.GetSubjectRegex()); err != nil {
return trace.BadParameter("spec.keyless.identities[%d].subject_regex: failed to parse regex: %v", idx, err)
}
default:
return trace.BadParameter("spec.keyless.identities[%d].subject_matcher: subject or subject_regex is required", idx)
}
}
roots := make(root.TrustedMaterialCollection, 0)
for idx, trustedRoot := range s.GetSpec().GetKeyless().GetTrustedRoots() {
root, err := root.NewTrustedRootFromJSON([]byte(trustedRoot))
if err != nil {
return trace.BadParameter("spec.keyless.trusted_roots[%d]: failed to parse trusted root: %v", idx, err)
}
roots = append(roots, root)
}
// If the user is overriding the default (Public Good Instance) trusted
// roots with their own, they must specify at least one transparency log
// or timestamp authority that can be used to verify keyless certificates.
if len(roots) != 0 && len(roots.CTLogs()) == 0 && len(roots.TimestampingAuthorities()) == 0 {
return trace.BadParameter("spec.keyless.trusted_roots: must configure at least one transparency log or timestamp authority")
}
}
requirements := s.GetSpec().GetRequirements()
if requirements.GetArtifactSignature() && len(requirements.GetAttestations()) != 0 {
return trace.BadParameter("spec.requirements: artifact_signature and attestations are mutually exclusive")
}
if !requirements.GetArtifactSignature() && len(requirements.GetAttestations()) == 0 {
return trace.BadParameter("spec.requirements: either artifact_signature or attestations is required")
}
for idx, attestation := range requirements.GetAttestations() {
if attestation.GetPredicateType() == "" {
return trace.BadParameter("spec.requirements.attestations[%d].predicate_type: is required", idx)
}
}
return nil
}
// Teleport
// Copyright (C) 2024 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"
"net/url"
"strings"
"github.com/gravitational/trace"
"github.com/spiffe/go-spiffe/v2/bundle/spiffebundle"
"github.com/spiffe/go-spiffe/v2/spiffeid"
machineidv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
"github.com/gravitational/teleport/api/types"
)
// SPIFFEFederations is an interface over the SPIFFEFederations service. This
// interface may also be implemented by a client to allow remote and local
// consumers to access the resource in a similar way.
type SPIFFEFederations interface {
// GetSPIFFEFederation gets a SPIFFE Federation by name.
GetSPIFFEFederation(
ctx context.Context, name string,
) (*machineidv1.SPIFFEFederation, error)
// ListSPIFFEFederations lists all SPIFFE Federations using Google style
// pagination.
ListSPIFFEFederations(
ctx context.Context, pageSize int, lastToken string,
) ([]*machineidv1.SPIFFEFederation, string, error)
// CreateSPIFFEFederation creates a new SPIFFE Federation.
CreateSPIFFEFederation(
ctx context.Context, spiffeFederation *machineidv1.SPIFFEFederation,
) (*machineidv1.SPIFFEFederation, error)
// DeleteSPIFFEFederation deletes a SPIFFE Federation by name.
DeleteSPIFFEFederation(ctx context.Context, name string) error
// UpdateSPIFFEFederation updates a SPIFFE Federation. It will not act if the resource is not found
// or where the revision does not match.
UpdateSPIFFEFederation(
ctx context.Context, spiffeFederation *machineidv1.SPIFFEFederation,
) (*machineidv1.SPIFFEFederation, error)
}
// MarshalSPIFFEFederation marshals the SPIFFEFederation object into a JSON byte
// array.
func MarshalSPIFFEFederation(object *machineidv1.SPIFFEFederation, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalSPIFFEFederation unmarshals the SPIFFEFederation object from a
// JSON byte array.
func UnmarshalSPIFFEFederation(
data []byte, opts ...MarshalOption,
) (*machineidv1.SPIFFEFederation, error) {
return UnmarshalProtoResource[*machineidv1.SPIFFEFederation](data, opts...)
}
// ValidateSPIFFEFederation validates the SPIFFEFederation object.
func ValidateSPIFFEFederation(s *machineidv1.SPIFFEFederation) error {
switch {
case s == nil:
return trace.BadParameter("object cannot be nil")
case s.GetVersion() != types.V1:
return trace.BadParameter("version: only %q is supported", types.V1)
case s.GetKind() != types.KindSPIFFEFederation:
return trace.BadParameter("kind: must be %q", types.KindSPIFFEFederation)
case !s.HasMetadata():
return trace.BadParameter("metadata: is required")
case s.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name: is required")
case !s.HasSpec():
return trace.BadParameter("spec: is required")
case !s.GetSpec().HasBundleSource():
return trace.BadParameter("spec.bundle_source: is required")
case s.GetSpec().GetBundleSource().HasHttpsWeb() && s.GetSpec().GetBundleSource().HasStatic():
return trace.BadParameter("spec.bundle_source: at most one of https_web or static can be set")
case !s.GetSpec().GetBundleSource().HasHttpsWeb() && !s.GetSpec().GetBundleSource().HasStatic():
return trace.BadParameter("spec.bundle_source: at least one of https_web or static must be set")
}
// Validate name is valid SPIFFE Trust Domain name without the "spiffe://"
name := s.GetMetadata().GetName()
if strings.HasPrefix(name, "spiffe://") {
return trace.BadParameter(
"metadata.name: must not include the spiffe:// prefix",
)
}
td, err := spiffeid.TrustDomainFromString(name)
if err != nil {
return trace.Wrap(err, "validating metadata.name")
}
// Validate Static
if s.GetSpec().GetBundleSource().HasStatic() {
if s.GetSpec().GetBundleSource().GetStatic().GetBundle() == "" {
return trace.BadParameter("spec.bundle_source.static.bundle: is required")
}
// Validate contents
// TODO(noah): Is this a bit intense to run on every validation?
// This could easily be moved into reconciliation...
_, err := spiffebundle.Parse(td, []byte(s.GetSpec().GetBundleSource().GetStatic().GetBundle()))
if err != nil {
return trace.Wrap(err, "validating spec.bundle_source.static.bundle")
}
}
// Validate HTTPSWeb
if s.GetSpec().GetBundleSource().HasHttpsWeb() {
if s.GetSpec().GetBundleSource().GetHttpsWeb().GetBundleEndpointUrl() == "" {
return trace.BadParameter("spec.bundle_source.https_web.bundle_endpoint_url: is required")
}
_, err := url.Parse(s.GetSpec().GetBundleSource().GetHttpsWeb().GetBundleEndpointUrl())
if err != nil {
return trace.Wrap(err, "validating spec.bundle_source.https_web.bundle_endpoint_url")
}
}
// Ensure that all key status fields are set if any are set. This is a safeguard against weird inconsistent states
// where some fields are set and others are not.
currentBundleSet := s.GetStatus().GetCurrentBundle() != ""
currentBundledSyncedAtSet := s.GetStatus().GetCurrentBundleSyncedAt() != nil
currentBundleSyncedFromSet := s.GetStatus().GetCurrentBundleSyncedFrom() != nil
anyStatusFieldSet := currentBundleSet || currentBundledSyncedAtSet || currentBundleSyncedFromSet
allStatusFieldsSet := currentBundleSet && currentBundledSyncedAtSet && currentBundleSyncedFromSet
if anyStatusFieldSet && !allStatusFieldsSet {
return trace.BadParameter("status: all of ['current_bundle', 'current_bundle_synced_at', 'current_bundle_synced_from'] must be set if any are set")
}
return nil
}
/*
* Teleport
* Copyright (C) 2026 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 (
"time"
"github.com/gravitational/trace"
"google.golang.org/protobuf/proto"
"github.com/gravitational/teleport/api/constants"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
"github.com/gravitational/teleport/api/types"
)
// SSHAccessChecker provides SSH-specific access checking, abstracting over scoped and unscoped identities.
// It is obtained from [ScopedAccessChecker.SSH] and should not be constructed directly. Methods on this type
// implement SSH-specific behavior, branching internally between the scoped and unscoped paths of the underlying
// [ScopedAccessChecker].
type SSHAccessChecker struct {
checker *ScopedAccessChecker
}
// CheckAccessToSSHServer checks access to an SSH server for the given OS user.
func (c *SSHAccessChecker) CheckAccessToSSHServer(target types.Server, state AccessState, osUser string) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAccess(target, state, NewLoginMatcher(osUser))
}
return c.checker.scopedCompatChecker.CheckAccess(target, state, NewLoginMatcher(osUser))
}
// CanAccessSSHServer checks whether read access to the specified SSH server is possible without
// regard to a specific OS user or MFA state. Used for listing/filtering.
func (c *SSHAccessChecker) CanAccessSSHServer(target types.Server) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAccess(target, AccessState{MFAVerified: true})
}
return c.checker.scopedCompatChecker.CheckAccess(target, AccessState{MFAVerified: true})
}
// AdjustClientIdleTimeout determines the SSH client idle timeout to apply. The supplied argument must be
// the globally defined most-permissive value. For scoped identities, the value is read directly from the
// scoped role proto (ssh.client_idle_timeout takes precedence over defaults.client_idle_timeout). If the
// role specifies a more restrictive value it is returned; otherwise the global value is returned unchanged.
// An error is returned if the role contains a non-empty duration string that cannot be parsed.
func (c *SSHAccessChecker) AdjustClientIdleTimeout(timeout time.Duration) (time.Duration, error) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.AdjustClientIdleTimeout(timeout), nil
}
return c.checker.adjustScopedClientIdleTimeout(c.checker.role.GetSpec().GetSsh().GetClientIdleTimeout(), timeout)
}
// AdjustDisconnectExpiredCert adjusts whether to disconnect on certificate expiry.
func (c *SSHAccessChecker) AdjustDisconnectExpiredCert(disconnect bool) bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.AdjustDisconnectExpiredCert(disconnect)
}
ssh := c.checker.role.GetSpec().GetSsh()
var disconnectExpiredCert *bool
if ssh != nil {
disconnectExpiredCert = proto.ValueOrNil(ssh.HasDisconnectExpiredCert(), ssh.GetDisconnectExpiredCert)
}
return c.checker.adjustScopedDisconnectExpiredCert(disconnectExpiredCert, disconnect)
}
// LockingMode returns the SSH lock enforcement mode to apply.
func (c *SSHAccessChecker) LockingMode(defaultMode constants.LockingMode) constants.LockingMode {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.LockingMode(defaultMode)
}
return c.checker.scopedLockingMode(c.checker.role.GetSpec().GetSsh().GetLock(), defaultMode)
}
// SessionRecordingMode returns the session recording mode for SSH sessions.
// SSH recording mode takes precedence over the defaults
func (c *SSHAccessChecker) SessionRecordingMode() constants.SessionRecordingMode {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.SessionRecordingMode(constants.SessionRecordingServiceSSH)
}
sr := c.checker.role.GetSpec().GetSsh().GetSessionRecording()
if sr == nil {
sr = c.checker.role.GetSpec().GetDefaults().GetSessionRecording()
}
if sr.GetMode() != "" {
return constants.SessionRecordingMode(sr.GetMode())
}
return constants.SessionRecordingModeBestEffort
}
// CanPortForward returns true if port forwarding is permitted.
func (c *SSHAccessChecker) CanPortForward() bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CanPortForward()
}
return c.checker.scopedCompatChecker.CanPortForward()
}
// CanForwardAgents returns true if SSH agent forwarding is permitted.
func (c *SSHAccessChecker) CanForwardAgents() bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CanForwardAgents()
}
return c.checker.scopedCompatChecker.CanForwardAgents()
}
// PermitX11Forwarding returns true if X11 forwarding is permitted.
func (c *SSHAccessChecker) PermitX11Forwarding() bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.PermitX11Forwarding()
}
return c.checker.role.GetSpec().GetSsh().GetPermitX11Forwarding()
}
// SSHPortForwardMode returns the SSH port forwarding mode.
func (c *SSHAccessChecker) SSHPortForwardMode() decisionpb.SSHPortForwardMode {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.SSHPortForwardMode()
}
remote := c.checker.role.GetSpec().GetSsh().GetPortForwarding().GetRemote()
local := c.checker.role.GetSpec().GetSsh().GetPortForwarding().GetLocal()
var denyRemote, denyLocal bool
if remote != nil && remote.HasEnabled() && !remote.GetEnabled() {
denyRemote = true
}
if local != nil && local.HasEnabled() && !local.GetEnabled() {
denyLocal = true
}
// enforcing implicit allow and preferring allow over explicit deny
switch {
case denyRemote && denyLocal:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_OFF
case denyRemote:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_LOCAL
case denyLocal:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_REMOTE
default:
return decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_ON
}
}
// HostSudoers returns the sudoers rules for the host.
func (c *SSHAccessChecker) HostSudoers(srv types.Server) ([]string, error) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.HostSudoers(srv)
}
return c.checker.role.GetSpec().GetSsh().GetHostSudoers(), nil
}
// EnhancedRecordingSet returns the set of enhanced session recording events to capture.
func (c *SSHAccessChecker) EnhancedRecordingSet() map[string]bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.EnhancedRecordingSet()
}
events := c.checker.role.GetSpec().GetSsh().GetEnhancedRecording()
m := make(map[string]bool)
if events.GetCommand() {
m[constants.EnhancedRecordingCommand] = true
}
if events.GetDisk() {
m[constants.EnhancedRecordingDisk] = true
}
if events.GetNetwork() {
m[constants.EnhancedRecordingNetwork] = true
}
return m
}
// HostUsers returns host user creation information for the server, or nil if host user creation is disabled.
func (c *SSHAccessChecker) HostUsers(srv types.Server) (*HostUsersDecision, error) {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.HostUsers(srv)
}
createHostUser := c.checker.role.GetSpec().GetSsh().GetHostUserCreation()
var hostUserMode types.CreateHostUserMode
if err := hostUserMode.UnmarshalText([]byte(createHostUser.GetMode())); err != nil {
return nil, trace.Wrap(err)
}
// If no create_host_user block, or mode is OFF/UNSPECIFIED, host user creation is disabled.
if hostUserMode == types.CreateHostUserMode_HOST_USER_MODE_OFF ||
hostUserMode == types.CreateHostUserMode_HOST_USER_MODE_UNSPECIFIED {
return &HostUsersDecision{
Info: nil,
DeniedBy: []*decisionpb.Determinant{decisionpb.Determinant_builder{
Kind: c.checker.role.GetKind(),
Name: c.checker.role.GetMetadata().GetName(),
}.Build()},
}, nil
}
// Convert to decision
var decisionMode decisionpb.HostUserMode
switch hostUserMode {
case types.CreateHostUserMode_HOST_USER_MODE_KEEP:
decisionMode = decisionpb.HostUserMode_HOST_USER_MODE_KEEP
case types.CreateHostUserMode_HOST_USER_MODE_INSECURE_DROP:
decisionMode = decisionpb.HostUserMode_HOST_USER_MODE_DROP
default:
decisionMode = decisionpb.HostUserMode_HOST_USER_MODE_UNSPECIFIED
}
traits := c.checker.Traits()
var uid, gid string
if uidL := traits[constants.TraitHostUserUID]; len(uidL) >= 1 {
uid = uidL[0]
}
if gidL := traits[constants.TraitHostUserGID]; len(gidL) >= 1 {
gid = gidL[0]
}
return &HostUsersDecision{
Info: decisionpb.HostUsersInfo_builder{
Groups: createHostUser.GetGroups(),
Mode: decisionMode,
Uid: uid,
Gid: gid,
Shell: createHostUser.GetShell(),
}.Build(),
AllowedBy: []*decisionpb.Determinant{decisionpb.Determinant_builder{
Kind: c.checker.role.GetKind(),
Name: c.checker.role.GetMetadata().GetName(),
}.Build()},
}, nil
}
// CheckAgentForward checks whether SSH agent forwarding is permitted for the given login.
func (c *SSHAccessChecker) CheckAgentForward(login string) error {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CheckAgentForward(login)
}
if !c.checker.role.GetSpec().GetSsh().GetForwardAgent() {
return trace.AccessDenied("agent forwarding is not permitted for scoped role %q",
c.checker.role.GetMetadata().GetName())
}
return nil
}
// MaxConnections returns the maximum number of concurrent SSH connections permitted.
// A value of zero means unconstrained.
func (c *SSHAccessChecker) MaxConnections() int64 {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.MaxConnections()
}
return c.checker.scopedCompatChecker.MaxConnections()
}
// MaxSessions returns the maximum number of concurrent SSH sessions per connection permitted.
// A value of zero means unconstrained.
func (c *SSHAccessChecker) MaxSessions() int64 {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.MaxSessions()
}
return c.checker.role.GetSpec().GetSsh().GetMaxSessions()
}
// getScopedLogins returns the OS logins permitted by this scoped role. Returns nil for unscoped
// identities, which aggregate logins differently via [CertificateParameterContext.GetSSHLoginsForTTL].
// This method is intentionally unexported to prevent accidental use outside cert-param aggregation.
func (c *SSHAccessChecker) getScopedLogins() []string {
if !c.checker.isScoped() {
return nil
}
return c.checker.role.GetSpec().GetSsh().GetLogins()
}
// CanCopyFiles returns true if remote file operations via SCP or SFTP are permitted.
// If GetSshFileCopy is nil, then we default to true.
func (c *SSHAccessChecker) CanCopyFiles() bool {
if !c.checker.isScoped() {
return c.checker.unscopedChecker.CanCopyFiles()
}
ssh := c.checker.role.GetSpec().GetSsh()
if ssh == nil || !ssh.HasFileCopy() {
return true
}
return ssh.GetFileCopy()
}
/*
* Teleport
* Copyright (C) 2024 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/trace"
userprovisioningpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/userprovisioning/v2"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/services/label"
)
// StaticHostUserService manages host users that should be created on SSH nodes.
type StaticHostUser interface {
// ListStaticHostUsers lists static host users.
ListStaticHostUsers(ctx context.Context, pageSize int, pageToken string) ([]*userprovisioningpb.StaticHostUser, string, error)
// GetStaticHostUser returns a static host user by name.
GetStaticHostUser(ctx context.Context, name string) (*userprovisioningpb.StaticHostUser, error)
// CreateStaticHostUser creates a static host user.
CreateStaticHostUser(ctx context.Context, in *userprovisioningpb.StaticHostUser) (*userprovisioningpb.StaticHostUser, error)
// UpdateStaticHostUser updates a static host user.
UpdateStaticHostUser(ctx context.Context, in *userprovisioningpb.StaticHostUser) (*userprovisioningpb.StaticHostUser, error)
// UpsertStaticHostUser upserts a static host user.
UpsertStaticHostUser(ctx context.Context, in *userprovisioningpb.StaticHostUser) (*userprovisioningpb.StaticHostUser, error)
// DeleteStaticHostUser deletes a static host user. Note that this does not
// remove any host users created on nodes from the resource.
DeleteStaticHostUser(ctx context.Context, name string) error
}
// ValidateStaticHostUser checks that required parameters are set for the
// specified StaticHostUser.
func ValidateStaticHostUser(u *userprovisioningpb.StaticHostUser) error {
// Check if required info exists.
if u == nil {
return trace.BadParameter("StaticHostUser is nil")
}
if !u.HasMetadata() {
return trace.BadParameter("Metadata is nil")
}
if u.GetMetadata().GetName() == "" {
return trace.BadParameter("missing name")
}
if !u.HasSpec() {
return trace.BadParameter("Spec is nil")
}
if len(u.GetSpec().GetMatchers()) == 0 {
return trace.BadParameter("missing matchers")
}
for _, matcher := range u.GetSpec().GetMatchers() {
// Check if matcher can match any resources.
if len(matcher.GetNodeLabels()) == 0 && len(matcher.GetNodeLabelsExpression()) == 0 {
return trace.BadParameter("either NodeLabels or NodeLabelsExpression must be set")
}
for _, label := range matcher.GetNodeLabels() {
if label.GetName() == types.Wildcard && (len(label.GetValues()) != 1 || label.GetValues()[0] != types.Wildcard) {
return trace.BadParameter("selector *:<val> is not supported")
}
}
if len(matcher.GetNodeLabelsExpression()) > 0 {
if _, err := label.ParseExpression(matcher.GetNodeLabelsExpression()); err != nil {
return trace.BadParameter("parsing node labels expression: %v", err)
}
}
}
return nil
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalStaticTokens unmarshals the StaticTokens resource from JSON.
func UnmarshalStaticTokens(bytes []byte, opts ...MarshalOption) (types.StaticTokens, error) {
var staticTokens types.StaticTokensV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if err := utils.FastUnmarshal(bytes, &staticTokens); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := staticTokens.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
staticTokens.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
staticTokens.SetExpiry(cfg.Expires)
}
return &staticTokens, nil
}
// MarshalStaticTokens marshals the StaticTokens resource to JSON.
func MarshalStaticTokens(staticToken types.StaticTokens, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch staticToken := staticToken.(type) {
case *types.StaticTokensV2:
if err := staticToken.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, staticToken))
default:
return nil, trace.BadParameter("unrecognized static token version %T", staticToken)
}
}
// Teleport
// Copyright (C) 2026 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"
subcav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/subca/v1"
"github.com/gravitational/teleport/api/types"
)
// SubCAServiceGetter is the read-only SubCAService interface.
//
// See lib/services/local.SubCAService.
type SubCAServiceGetter interface {
// GetCertAuthorityOverride reads a CA override resource by ID.
GetCertAuthorityOverride(ctx context.Context, id types.CertAuthorityOverrideID) (*subcav1.CertAuthorityOverride, error)
// ListCertAuthorityOverrides lists all CA overrides.
ListCertAuthorityOverrides(ctx context.Context, pageSize int, pageToken string) (_ []*subcav1.CertAuthorityOverride, nextPageToken string, _ error)
}
// SubCAService manages CertAuthorityOverride resources.
//
// See lib/services/local.SubCAService.
type SubCAService interface {
SubCAServiceGetter
// CreateCertAuthorityOverride creates a CA override.
CreateCertAuthorityOverride(ctx context.Context, resource *subcav1.CertAuthorityOverride) (*subcav1.CertAuthorityOverride, error)
// DeleteCertAuthorityOverride hard-deletes a CA override.
DeleteCertAuthorityOverride(ctx context.Context, id types.CertAuthorityOverrideID) error
// ConditionalDeleteCertAuthorityOverride conditionally deletes a CA override
// based on its revision.
ConditionalDeleteCertAuthorityOverride(
ctx context.Context,
id types.CertAuthorityOverrideID,
revision string,
) error
// UpdateCertAuthorityOverride conditionally updates a CA override.
UpdateCertAuthorityOverride(ctx context.Context, resource *subcav1.CertAuthorityOverride) (*subcav1.CertAuthorityOverride, error)
// UpsertCertAuthorityOverride unconditionally creates or updates a CA override.
UpsertCertAuthorityOverride(ctx context.Context, resource *subcav1.CertAuthorityOverride) (*subcav1.CertAuthorityOverride, error)
}
// MarshalCertAuthorityOverride marshals a CA override resource.
func MarshalCertAuthorityOverride(resource *subcav1.CertAuthorityOverride, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(resource, opts...)
}
// UnmarshalCertAuthorityOverride unmarshals a CA override resource.
func UnmarshalCertAuthorityOverride(data []byte, opts ...MarshalOption) (*subcav1.CertAuthorityOverride, error) {
return UnmarshalProtoResource[*subcav1.CertAuthorityOverride](data, opts...)
}
// PendingCSRRequestServiceGetter is the read-only PendingCSRRequestService
// interface.
//
// This service is not exposed to Auth clients.
type PendingCSRRequestServiceGetter interface {
// GetPendingCSRRequest reads a PendingCSRRequest by name.
GetPendingCSRRequest(ctx context.Context, name string) (*subcav1.PendingCSRRequest, error)
// ListPendingCSRRequests lists all PendingCSRRequests.
ListPendingCSRRequests(ctx context.Context, pageSize int, pageToken string) (_ []*subcav1.PendingCSRRequest, nextPageToken string, _ error)
}
// PendingCSRRequestService manages PendingCSRRequest resources.
//
// This service is not exposed to Auth clients.
type PendingCSRRequestService interface {
PendingCSRRequestServiceGetter
// CreatePendingCSRRequest creates a PendingCSRRequest.
//
// PendingCSRRequest instances must have an expiration. If they don't a
// default expiration is assigned on creation.
CreatePendingCSRRequest(ctx context.Context, resource *subcav1.PendingCSRRequest) (*subcav1.PendingCSRRequest, error)
// UpdatePendingCSRRequest conditionally updates a PendingCSRRequest.
UpdatePendingCSRRequest(ctx context.Context, resource *subcav1.PendingCSRRequest) (*subcav1.PendingCSRRequest, error)
// DeletePendingCSRRequest hard-deletes a PendingCSRRequest.
DeletePendingCSRRequest(ctx context.Context, name string) error
}
// Teleport
// Copyright (C) 2025 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"
"iter"
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/vulcand/predicate"
"github.com/gravitational/teleport"
summarizerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/summarizer/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/events"
apisummarizer "github.com/gravitational/teleport/api/types/summarizer"
)
// Summarizer is a service that provides methods to manage summary inference
// configuration resources in the backend.
type Summarizer interface {
SummarizerServiceGetter
// CreateInferenceModel creates a new session summary inference model in the
// backend.
CreateInferenceModel(ctx context.Context, model *summarizerv1.InferenceModel) (*summarizerv1.InferenceModel, error)
// UpdateInferenceModel updates an existing session summary inference model
// in the backend.
UpdateInferenceModel(ctx context.Context, model *summarizerv1.InferenceModel) (*summarizerv1.InferenceModel, error)
// UpsertInferenceModel creates or updates a session summary inference model
// in the backend. If the model already exists, it will be updated.
UpsertInferenceModel(ctx context.Context, model *summarizerv1.InferenceModel) (*summarizerv1.InferenceModel, error)
// DeleteInferenceModel deletes a session summary inference model from the
// backend by name.
DeleteInferenceModel(ctx context.Context, name string) error
// CreateInferenceSecret creates a new session summary inference secret in
// the backend. The returned object contains the secret value and should be
// handled with care.
CreateInferenceSecret(ctx context.Context, secret *summarizerv1.InferenceSecret) (*summarizerv1.InferenceSecret, error)
// UpdateInferenceSecret updates an existing session summary inference secret
// in the backend. The returned object contains the secret value and should
// be handled with care.
UpdateInferenceSecret(ctx context.Context, secret *summarizerv1.InferenceSecret) (*summarizerv1.InferenceSecret, error)
// UpsertInferenceSecret creates or updates a session summary inference
// secretin the backend. If the secret already exists, it will be updated.
// The returned object contains the secret value and should be handled with
// care.
UpsertInferenceSecret(ctx context.Context, secret *summarizerv1.InferenceSecret) (*summarizerv1.InferenceSecret, error)
// DeleteInferenceSecret deletes a session summary inference secret from the
// backend by name.
DeleteInferenceSecret(ctx context.Context, name string) error
// CreateInferencePolicy creates a new session summary inference policy in
// the backend.
CreateInferencePolicy(ctx context.Context, policy *summarizerv1.InferencePolicy) (*summarizerv1.InferencePolicy, error)
// UpdateInferencePolicy updates an existing session summary inference policy
// in the backend.
UpdateInferencePolicy(ctx context.Context, policy *summarizerv1.InferencePolicy) (*summarizerv1.InferencePolicy, error)
// UpsertInferencePolicy creates or updates a session summary inference
// policy in the backend. If the policy already exists, it will be updated.
UpsertInferencePolicy(ctx context.Context, policy *summarizerv1.InferencePolicy) (*summarizerv1.InferencePolicy, error)
// DeleteInferencePolicy deletes a session summary inference policy from the
// backend by name.
DeleteInferencePolicy(ctx context.Context, name string) error
// CreateClassifier creates a new session summarization classifier in the
// backend.
CreateClassifier(ctx context.Context, classifier *summarizerv1.Classifier) (*summarizerv1.Classifier, error)
// UpdateClassifier updates an existing session summarization classifier in
// the backend.
UpdateClassifier(ctx context.Context, classifier *summarizerv1.Classifier) (*summarizerv1.Classifier, error)
// UpsertClassifier creates or updates a session summarization classifier
// in the backend. If the classifier already exists, it will be updated.
UpsertClassifier(ctx context.Context, classifier *summarizerv1.Classifier) (*summarizerv1.Classifier, error)
// DeleteClassifier deletes a session summarization classifier from the
// backend by name.
DeleteClassifier(ctx context.Context, name string) error
// CreateRetrievalModel creates the search model in the backend.
// Only one RetrievalModel can exist per cluster.
CreateRetrievalModel(ctx context.Context, model *summarizerv1.RetrievalModel) (*summarizerv1.RetrievalModel, error)
// UpdateRetrievalModel updates the existing search model in the backend.
UpdateRetrievalModel(ctx context.Context, model *summarizerv1.RetrievalModel) (*summarizerv1.RetrievalModel, error)
// UpsertRetrievalModel creates or updates the search model in the backend.
// If the model already exists, it will be updated.
UpsertRetrievalModel(ctx context.Context, model *summarizerv1.RetrievalModel) (*summarizerv1.RetrievalModel, error)
// DeleteRetrievalModel deletes the search model from the backend.
// Since only one RetrievalModel can exist per cluster, no name is required.
DeleteRetrievalModel(ctx context.Context) error
}
// SummarizerServiceGetter is the interface that defines the methods required for
// retrieving objects from the cache.
type SummarizerServiceGetter interface {
// GetInferenceModel retrieves a session summary inference model from the
// backend by name.
GetInferenceModel(ctx context.Context, name string) (*summarizerv1.InferenceModel, error)
// ListInferenceModels lists session summary inference models in the backend
// with pagination support. Returns a slice of models and a next page token.
ListInferenceModels(ctx context.Context, size int, pageToken string) ([]*summarizerv1.InferenceModel, string, error)
// GetInferenceSecret retrieves a session summary inference secret from the
// backend by name. The returned object contains the secret value and should
// be handled with care.
GetInferenceSecret(ctx context.Context, name string) (*summarizerv1.InferenceSecret, error)
// ListInferenceSecrets lists session summary inference secrets in the
// backend with pagination support. Returns a slice of secrets and a next
// page token. The returned objects contain the secret values and should be
// handled with care.
ListInferenceSecrets(ctx context.Context, size int, pageToken string) ([]*summarizerv1.InferenceSecret, string, error)
// GetInferencePolicy retrieves a session summary inference policy from the
// backend by name.
GetInferencePolicy(ctx context.Context, name string) (*summarizerv1.InferencePolicy, error)
// ListInferencePolicies lists session summary inference policies in the
// backend with pagination support. Returns a slice of policies and a next
// page token.
ListInferencePolicies(ctx context.Context, size int, pageToken string) ([]*summarizerv1.InferencePolicy, string, error)
// AllInferencePolicies returns an iterator that retrieves all session
// summary inference policies from the backend, without pagination.
AllInferencePolicies(ctx context.Context) iter.Seq2[*summarizerv1.InferencePolicy, error]
// GetClassifier retrieves a session summarization classifier from the
// backend by name.
GetClassifier(ctx context.Context, name string) (*summarizerv1.Classifier, error)
// ListClassifiers lists session summarization classifiers in the backend
// with pagination support. Returns a slice of classifiers and a next page
// token.
ListClassifiers(ctx context.Context, size int, pageToken string) ([]*summarizerv1.Classifier, string, error)
// RangeClassifiers returns an iterator that retrieves session summarization
// classifiers from the backend, without pagination, starting with the
// resource named start and ending before the resource named end. Empty
// bounds iterate from the beginning and/or to the end of the collection.
RangeClassifiers(ctx context.Context, start, end string) iter.Seq2[*summarizerv1.Classifier, error]
// GetRetrievalModel retrieves the search model from the backend.
// Since only one RetrievalModel can exist per cluster, no name is required.
GetRetrievalModel(ctx context.Context) (*summarizerv1.RetrievalModel, error)
}
// InferencePolicyMatchingContext is a special kind of [RuleContext] that is
// used for matching inference policies to sessions using predicates. It also
// allows validating inference policy filter expressions, since it matches
// identifiers for any supported resource and session event types, regardless
// which one is being used (or if none is).
type InferencePolicyMatchingContext struct {
// User is the user who initiated the session.
User UserState
// Resource is the resource being accessed.
Resource types.Resource
// Session is a session.end or windows.desktop.session.end event. These
// events hold information about session recordings.
Session events.AuditEvent
}
// GetIdentifier returns the value of an identifier defined in a context.
func (ctx *InferencePolicyMatchingContext) GetIdentifier(fields []string) (any, error) {
switch fields[0] {
case UserIdentifier:
var user UserState
if ctx.User == nil {
user = emptyUser
} else {
user = ctx.User
}
val, err := predicate.GetFieldByTag(user, teleport.JSON, fields[1:])
return val, trace.Wrap(err)
case ResourceIdentifier:
// First, try to fetch field value from the resource in the context.
val, origErr := predicate.GetFieldByTag(ctx.Resource, teleport.JSON, fields[1:])
if origErr == nil {
return val, nil
}
if !trace.IsNotFound(origErr) {
return nil, trace.Wrap(origErr)
}
// Otherwise, try to fetch field value from dummy resources of all
// supported types to figure out if it exists in any of the supported
// types. If it does, a zero value is returned; otherwise, an error is
// returned.
for _, dummyResource := range []types.Resource{
&types.ServerV2{}, &types.KubernetesClusterV3{}, &types.DatabaseV3{},
&types.WindowsDesktopV3{},
} {
zeroVal, err := predicate.GetFieldByTag(dummyResource, teleport.JSON, fields[1:])
if err == nil {
return zeroVal, nil
}
if trace.IsNotFound(err) {
continue
}
return val, trace.Wrap(origErr)
}
return val, trace.Wrap(origErr)
case SessionIdentifier:
// First, try to fetch field value from the session in the context.
var session events.AuditEvent = &events.SessionEnd{}
switch ctx.Session.(type) {
case *events.SessionEnd, *events.DatabaseSessionEnd, *events.WindowsDesktopSessionEnd:
session = ctx.Session
}
val, origErr := predicate.GetFieldByTag(session, teleport.JSON, fields[1:])
if origErr == nil {
return val, nil
}
if !trace.IsNotFound(origErr) {
return nil, trace.Wrap(origErr)
}
// Otherwise, try to fetch field value from dummy events of all supported
// types to figure out if it exists in any of the supported types. If it
// does, a zero value is returned; otherwise, an error is returned.
if zeroVal, err := getMissingEmptyFieldForSessionEnd(fields); err == nil {
return zeroVal, nil
}
return val, trace.Wrap(origErr)
default:
return nil, trace.NotFound("%v is not defined", strings.Join(fields, "."))
}
}
// Returns an error, since this context does not support access checks.
func (ctx *InferencePolicyMatchingContext) GetAccessChecker() (AccessChecker, error) {
return nil, trace.NotFound(
"access checker is not supported by InferencePolicyMatchingContext",
)
}
// GetResource returns resource specified in the context,
// returns error if not specified.
func (ctx *InferencePolicyMatchingContext) GetResource() (types.Resource, error) {
if ctx.Resource == nil {
return nil, trace.NotFound("resource is not set in the context")
}
return ctx.Resource, nil
}
// ExtendWithSessionEnd extends the context with a session end event and
// rebuilds the resource from the event.
func (ctx *InferencePolicyMatchingContext) ExtendWithSessionEnd(sessionEnd events.AuditEvent) {
ctx.Session = sessionEnd
ctx.Resource = rebuildResourceFromSessionEndEvent(sessionEnd)
}
// MatchingClassifiers returns the classifiers from the given sequence that
// apply to a session of the given kind and matching context. Classifiers are
// matched by session kind and filter expression the same way inference
// policies are, except that all matching classifiers are returned rather than
// the first one. The sequence is typically
// [SummarizerServiceGetter.RangeClassifiers].
func MatchingClassifiers(
classifiers iter.Seq2[*summarizerv1.Classifier, error],
sessionKind types.SessionKind,
matchingCtx *InferencePolicyMatchingContext,
) ([]*summarizerv1.Classifier, error) {
parser, err := NewWhereParser(matchingCtx)
if err != nil {
return nil, trace.Wrap(err)
}
var matched []*summarizerv1.Classifier
for c, err := range classifiers {
if err != nil {
return nil, trace.Wrap(err)
}
if !slices.Contains(c.GetSpec().GetKinds(), string(sessionKind)) {
continue
}
if filter := c.GetSpec().GetFilter(); filter != "" {
parseResult, err := parser.Parse(filter)
if err != nil {
return nil, trace.Wrap(err)
}
pred, ok := parseResult.(predicate.BoolPredicate)
if !ok {
return nil, trace.BadParameter("unsupported type: %T", parseResult)
}
if !pred() {
continue
}
}
matched = append(matched, c)
}
return matched, nil
}
// ValidateClassifier validates a classifier, including checking filter
// syntax. This function wraps [apisummarizer.ValidateClassifier], as no
// function in the api/types tree can depend on the lib/services package.
func ValidateClassifier(c *summarizerv1.Classifier) error {
err := apisummarizer.ValidateClassifier(c)
if err != nil {
return trace.Wrap(err)
}
s := c.GetSpec()
if s.GetFilter() != "" {
parser, err := NewWhereParser(&InferencePolicyMatchingContext{})
if err != nil {
return trace.Wrap(err)
}
parseResult, err := parser.Parse(s.GetFilter())
if err != nil {
return trace.Wrap(err, "spec.filter has to be a valid predicate")
}
if _, ok := parseResult.(predicate.BoolPredicate); !ok {
return trace.BadParameter("spec.filter has to be a boolean expression")
}
}
return nil
}
// ValidateInferencePolicy validates an inference policy, including checking
// filter syntax. This function wraps [apisummarizer.ValidateInferencePolicy],
// as no function in the api/types tree can depend on the lib/services package.
func ValidateInferencePolicy(p *summarizerv1.InferencePolicy) error {
err := apisummarizer.ValidateInferencePolicy(p)
if err != nil {
return trace.Wrap(err)
}
s := p.GetSpec()
if s.GetFilter() != "" {
parser, err := NewWhereParser(&InferencePolicyMatchingContext{})
if err != nil {
return trace.Wrap(err)
}
if _, err = parser.Parse(s.GetFilter()); err != nil {
return trace.Wrap(err, "spec.filter has to be a valid predicate")
}
}
return nil
}
/*
* 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"
"fmt"
"log/slog"
"regexp"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/parse"
)
// maxMismatchedTraitValuesLogged indicates the maximum number of trait values (that do not match a
// certain expression) to be shown in the log
const maxMismatchedTraitValuesLogged = 100
// TraitsToRoles maps the supplied traits to a list of teleport role names.
// Returns the list of roles mapped from traits.
// `warnings` optionally contains the list of warnings potentially interesting to the user.
func TraitsToRoles(ms types.TraitMappingSet, traits map[string][]string) (warnings []string, roles []string) {
warnings = traitsToRoles(ms, traits, func(role string, expanded bool) {
roles = append(roles, role)
})
return warnings, apiutils.Deduplicate(roles)
}
// TraitsToRoleMatchers maps the supplied traits to a list of role matchers. Prefer calling
// this function directly rather than calling TraitsToRoles and then building matchers from
// the resulting list since this function forces any roles which include substitutions to
// be literal matchers.
func TraitsToRoleMatchers(ms types.TraitMappingSet, traits map[string][]string) ([]parse.Matcher, error) {
var matchers []parse.Matcher
var firstErr error
traitsToRoles(ms, traits, func(role string, expanded bool) {
if expanded || utils.ContainsExpansion(role) {
// mapping process included variable expansion; we therefore
// "escape" normal matcher syntax and look only for exact matches.
// (note: this isn't about combatting maliciously constructed traits,
// traits are from trusted identity sources, this is just
// about avoiding unnecessary footguns).
matchers = append(matchers, literalMatcher{
value: role,
})
return
}
m, err := parse.NewMatcher(role)
if err != nil {
if firstErr == nil {
firstErr = err
}
return
}
matchers = append(matchers, m)
})
if firstErr != nil {
return nil, trace.Wrap(firstErr)
}
return matchers, nil
}
// traitsToRoles maps the supplied traits to teleport role names and passes them to a collector.
func traitsToRoles(ms types.TraitMappingSet, traits map[string][]string, collect func(role string, expanded bool)) (warnings []string) {
TraitMappingLoop:
for _, mapping := range ms {
var regexpIgnoreCase *regexp.Regexp
var regexp *regexp.Regexp
for traitName, traitValues := range traits {
if traitName != mapping.Trait {
continue
}
var mismatched []string
TraitLoop:
for _, traitValue := range traitValues {
for _, role := range mapping.Roles {
// this ensures that the case-insensitive regexp is compiled at most once, and only if strictly needed;
// after this if, regexpIgnoreCase must be non-nil
if regexpIgnoreCase == nil {
var err error
regexpIgnoreCase, err = utils.RegexpWithConfig(mapping.Value, utils.RegexpConfig{IgnoreCase: true})
if err != nil {
warnings = append(warnings, fmt.Sprintf("case-insensitive expression %q is not a valid regexp", mapping.Value))
continue TraitMappingLoop
}
}
// Run the initial replacement case-insensitively. Doing so will filter out all literal non-matches
// but will match on case discrepancies. We do another case-sensitive match below to see if the
// case is different
outRole, err := utils.ReplaceRegexpWith(regexpIgnoreCase, role, traitValue)
switch {
case err != nil:
// this trait value clearly did not match, move on to another
mismatched = append(mismatched, traitValue)
continue TraitLoop
case outRole == "":
case outRole != "":
// this ensures that the case-sensitive regexp is compiled at most once, and only if strictly needed;
// after this if, regexp must be non-nil
if regexp == nil {
var err error
regexp, err = utils.RegexpWithConfig(mapping.Value, utils.RegexpConfig{})
if err != nil {
warnings = append(warnings, fmt.Sprintf("case-sensitive expression %q is not a valid regexp", mapping.Value))
continue TraitMappingLoop
}
}
// Run the replacement case-sensitively to see if it matches.
// If there's no match, the trait specifies a mapping which is case-sensitive;
// we should log a warning but return an error.
// See https://github.com/gravitational/teleport/issues/6016 for details.
if _, err := utils.ReplaceRegexpWith(regexp, role, traitValue); err != nil {
warnings = append(warnings, fmt.Sprintf("trait %q matches value %q case-insensitively and would have yielded %q role", traitValue, mapping.Value, outRole))
continue
}
// skip empty replacement or empty role
collect(outRole, outRole != role)
}
}
}
// show at most maxMismatchedTraitValuesLogged trait values to prevent huge log lines
switch l := len(mismatched); {
case l > maxMismatchedTraitValuesLogged:
slog.
DebugContext(context.Background(), "trait value(s) did not match (showing first %d values)",
"mismatch_count", len(mismatched),
"max_mismatch_logged", maxMismatchedTraitValuesLogged,
"expression", mapping.Value,
"values", mismatched[0:maxMismatchedTraitValuesLogged],
)
case l > 0:
slog.DebugContext(context.Background(), "trait value(s) did not match",
"mismatch_count", len(mismatched),
"expression", mapping.Value,
"values", mismatched,
)
}
}
}
return
}
// literalMatcher is used to "escape" values which are not allowed to
// take advantage of normal matcher syntax by limiting them to only
// literal matches.
type literalMatcher struct {
value string
}
func (m literalMatcher) Match(in string) bool { return m.value == in }
func literalMatchers(literals []string) []parse.Matcher {
matchers := make([]parse.Matcher, 0, len(literals))
for _, literal := range literals {
matchers = append(matchers, literalMatcher{literal})
}
return matchers
}
/*
* 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 (
"errors"
"fmt"
"maps"
"slices"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/set"
)
// ValidateTrustedCluster checks and sets Trusted Cluster defaults
func ValidateTrustedCluster(tc types.TrustedCluster, allowEmptyRolesOpts ...bool) error {
if err := CheckAndSetDefaults(tc); err != nil {
return trace.Wrap(err)
}
// DELETE IN (7.0)
// This flag is used to allow reading trusted clusters with no role map.
// This was possible in OSS before 6.0 release.
allowEmptyRoles := false
if len(allowEmptyRolesOpts) != 0 {
allowEmptyRoles = allowEmptyRolesOpts[0]
}
// we are not mentioning Roles parameter because we are deprecating it
if len(tc.GetRoles()) == 0 && len(tc.GetRoleMap()) == 0 {
if !allowEmptyRoles {
return trace.BadParameter("missing 'role_map' parameter")
}
}
if _, err := parseRoleMap(tc.GetRoleMap()); err != nil {
return trace.Wrap(err)
}
return nil
}
// RoleMapToString prints user friendly representation of role mapping
func RoleMapToString(r types.RoleMap) string {
values, err := parseRoleMap(r)
if err != nil {
return fmt.Sprintf("<failed to parse: %v", err)
}
if len(values) != 0 {
return fmt.Sprintf("%v", values)
}
return "<empty>"
}
func parseRoleMap(r types.RoleMap) (map[string][]string, error) {
directMatch := make(map[string][]string)
for i := range r {
roleMap := r[i]
if roleMap.Remote == "" {
return nil, trace.BadParameter("missing 'remote' parameter for role_map")
}
_, err := utils.ReplaceRegexp(roleMap.Remote, "", "")
if trace.IsBadParameter(err) {
return nil, trace.BadParameter("failed to parse 'remote' parameter for role_map: %v", err.Error())
}
if len(roleMap.Local) == 0 {
return nil, trace.BadParameter("missing 'local' parameter for 'role_map'")
}
for _, local := range roleMap.Local {
if local == "" {
return nil, trace.BadParameter("missing 'local' property of 'role_map' entry")
}
if local == types.Wildcard {
return nil, trace.BadParameter("wildcard value is not supported for 'local' property of 'role_map' entry")
}
}
_, ok := directMatch[roleMap.Remote]
if ok {
return nil, trace.BadParameter("remote role '%v' match is already specified", roleMap.Remote)
}
directMatch[roleMap.Remote] = roleMap.Local
}
return directMatch, nil
}
func mapRoles(remoteUserRoles []string, r types.RoleMap) (map[string]set.Set[string], error) {
// define a mapping from local roles to the possible remote roles that may
// have granted them access
index := make(map[string]set.Set[string])
addToIndex := func(localRole, remoteRole string) {
remoteRoles, ok := index[localRole]
if !ok {
index[localRole] = set.New[string](remoteRole)
return
}
remoteRoles.Add(remoteRole)
}
// Run the role mapping forwards over the remote roles, collecting which
// localRoles are granted by which remote roles
for _, mapping := range r {
expression := mapping.Remote
for _, remoteRole := range remoteUserRoles {
// never map default implicit role, it is always
// added by default
if remoteRole == constants.DefaultImplicitRole {
continue
}
for _, replacementRole := range mapping.Local {
replacement, err := utils.ReplaceRegexp(expression, replacementRole, remoteRole)
switch {
case err == nil:
// empty replacement can occur when $2 expand refers
// to non-existing capture group in match expression
if replacement == "" {
continue
}
addToIndex(replacement, remoteRole)
case errors.Is(err, utils.ErrReplaceRegexNotFound):
continue
default:
return nil, trace.Wrap(err)
}
}
}
}
return index, nil
}
// MapRoles maps remote roles to local roles
func MapRoles(r types.RoleMap, remoteRoles []string) ([]string, error) {
_, err := parseRoleMap(r)
if err != nil {
return nil, trace.Wrap(err)
}
// when no remote roles are specified, assume that
// there is a single empty remote role (that should match wildcards)
if len(remoteRoles) == 0 {
remoteRoles = []string{""}
}
index, err := mapRoles(remoteRoles, r)
if err != nil {
return nil, trace.Wrap(err)
}
return slices.Collect(maps.Keys(index)), nil
}
// UnmapRoles attempts to deduce what remote-cluster roles a user might have,
// given the set of local roles the user has.
func UnmapRoles(r types.RoleMap, remoteUserRoles, localRoles []string) ([]string, error) {
index, err := mapRoles(remoteUserRoles, r)
if err != nil {
return nil, trace.Wrap(err)
}
// collect the remote roles of interest
remoteRoleSet := set.New[string]()
for _, localRole := range localRoles {
rr, ok := index[localRole]
if !ok {
return nil, trace.BadParameter("not all local roles could be mapped to remote roles")
}
remoteRoleSet.Add(rr.Elements()...)
}
return remoteRoleSet.Elements(), nil
}
// UnmarshalTrustedCluster unmarshals the TrustedCluster resource from JSON.
func UnmarshalTrustedCluster(bytes []byte, opts ...MarshalOption) (types.TrustedCluster, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var trustedCluster types.TrustedClusterV2
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
if err := utils.FastUnmarshal(bytes, &trustedCluster); err != nil {
return nil, trace.BadParameter("%s", err)
}
// DELETE IN(7.0)
// temporarily allow to read trusted cluster with no role map
// until users migrate from 6.0 OSS that had no role map present
const allowEmptyRoleMap = true
if err = ValidateTrustedCluster(&trustedCluster, allowEmptyRoleMap); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
trustedCluster.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
trustedCluster.SetExpiry(cfg.Expires)
}
return &trustedCluster, nil
}
// MarshalTrustedCluster marshals the TrustedCluster resource to JSON.
func MarshalTrustedCluster(trustedCluster types.TrustedCluster, opts ...MarshalOption) ([]byte, error) {
// DELETE IN(7.0)
// temporarily allow to read trusted cluster with no role map
// until users migrate from 6.0 OSS that had no role map present
const allowEmptyRoleMap = true
if err := ValidateTrustedCluster(trustedCluster, allowEmptyRoleMap); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch trustedCluster := trustedCluster.(type) {
case *types.TrustedClusterV2:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, trustedCluster))
default:
return nil, trace.BadParameter("unrecognized trusted cluster version %T", trustedCluster)
}
}
/*
* 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 (
"encoding/json"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// ValidateReverseTunnel validates the OIDC connector and sets default values
func ValidateReverseTunnel(rt types.ReverseTunnel) error {
if err := CheckAndSetDefaults(rt); err != nil {
return trace.Wrap(err)
}
for _, addr := range rt.GetDialAddrs() {
if _, err := utils.ParseAddr(addr); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// UnmarshalReverseTunnel unmarshals the ReverseTunnel resource from JSON.
func UnmarshalReverseTunnel(bytes []byte, opts ...MarshalOption) (types.ReverseTunnel, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing tunnel data")
}
var h types.ResourceHeader
err := json.Unmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V2:
var r types.ReverseTunnelV2
if err := utils.FastUnmarshal(bytes, &r); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := ValidateReverseTunnel(&r); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
r.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
r.SetExpiry(cfg.Expires)
}
return &r, nil
}
return nil, trace.BadParameter("reverse tunnel version %v is not supported", h.Version)
}
// MarshalReverseTunnel marshals the ReverseTunnel resource to JSON.
func MarshalReverseTunnel(reverseTunnel types.ReverseTunnel, opts ...MarshalOption) ([]byte, error) {
if err := ValidateReverseTunnel(reverseTunnel); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch reverseTunnel := reverseTunnel.(type) {
case *types.ReverseTunnelV2:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, reverseTunnel))
default:
return nil, trace.BadParameter("unrecognized reverse tunnel version %T", reverseTunnel)
}
}
/*
* 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 (
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// LatestTunnelConnection returns latest tunnel connection from the list
// of tunnel connections, if no connections found, returns NotFound error
func LatestTunnelConnection(conns []types.TunnelConnection) (types.TunnelConnection, error) {
var lastConn types.TunnelConnection
for i := range conns {
conn := conns[i]
if lastConn == nil || conn.GetLastHeartbeat().After(lastConn.GetLastHeartbeat()) {
lastConn = conn
}
}
if lastConn == nil {
return nil, trace.NotFound("no connections found")
}
return lastConn, nil
}
// TunnelConnectionStatus returns tunnel connection status based on the last
// heartbeat time recorded for a connection
func TunnelConnectionStatus(clock clockwork.Clock, conn types.TunnelConnection, offlineThreshold time.Duration) string {
diff := clock.Now().Sub(conn.GetLastHeartbeat())
if diff < offlineThreshold {
return teleport.RemoteClusterStatusOnline
}
return teleport.RemoteClusterStatusOffline
}
// UnmarshalTunnelConnection unmarshals TunnelConnection resource from JSON or YAML,
// sets defaults and checks the schema
func UnmarshalTunnelConnection(data []byte, opts ...MarshalOption) (types.TunnelConnection, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing tunnel connection data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
err = utils.FastUnmarshal(data, &h)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V2:
var r types.TunnelConnectionV2
if err := utils.FastUnmarshal(data, &r); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := r.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
r.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
r.SetExpiry(cfg.Expires)
}
return &r, nil
}
return nil, trace.BadParameter("reverse tunnel version %v is not supported", h.Version)
}
// MarshalTunnelConnection marshals the TunnelConnection resource to JSON.
func MarshalTunnelConnection(tunnelConnection types.TunnelConnection, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch tunnelConnection := tunnelConnection.(type) {
case *types.TunnelConnectionV2:
if err := tunnelConnection.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, tunnelConnection))
default:
return nil, trace.BadParameter("unrecognized tunnel connection version %T", tunnelConnection)
}
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalUIConfig unmarshals the UIConfig resource from JSON.
func UnmarshalUIConfig(data []byte, opts ...MarshalOption) (types.UIConfig, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing resource data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var uiconfig types.UIConfigV1
if err := utils.FastUnmarshal(data, &uiconfig); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := uiconfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
uiconfig.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
uiconfig.SetExpiry(cfg.Expires)
}
return &uiconfig, nil
}
// MarshalUIConfig marshals the UIConfig resource to JSON.
func MarshalUIConfig(uiconfig types.UIConfig, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch uiconfig := uiconfig.(type) {
case *types.UIConfigV1:
if err := uiconfig.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, uiconfig))
default:
return nil, trace.BadParameter("unrecognized uiconfig version %T", uiconfig)
}
}
/*
* 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"
"iter"
"log/slog"
"strings"
"sync"
"time"
"github.com/google/btree"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client/proto"
apidefaults "github.com/gravitational/teleport/api/defaults"
componentfeaturesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/componentfeatures/v1"
identitycenterv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/identitycenter/v1"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/componentfeatures"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// UnifiedResourceCacheConfig is used to configure a UnifiedResourceCache
type UnifiedResourceCacheConfig struct {
// BTreeDegree is a degree of B-Tree, 2 for example, will create a
// 2-3-4 tree (each node contains 1-3 items and 2-4 children).
BTreeDegree int
// Clock is a clock for time-related operations
Clock clockwork.Clock
// Component is a logging component
Component string
ResourceWatcherConfig
ResourceGetter
}
// UnifiedResourceCache contains a representation of all resources that are displayable in the UI
type UnifiedResourceCache struct {
rw sync.RWMutex
logger *slog.Logger
cfg UnifiedResourceCacheConfig
// nameTree is a BTree with items sorted by (hostname)/name/type
nameTree *btree.BTreeG[*item]
// typeTree is a BTree with items sorted by type/(hostname)/name
typeTree *btree.BTreeG[*item]
// resources is a map of all resources currently tracked in the tree
// the key is always name/type
resources map[string]resourceCollection
initializationC chan struct{}
stale bool
once sync.Once
cache *utils.FnCache
ResourceGetter
}
// NewUnifiedResourceCache creates a new memory cache that holds the unified resources
func NewUnifiedResourceCache(ctx context.Context, cfg UnifiedResourceCacheConfig) (*UnifiedResourceCache, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err, "setting defaults for unified resource cache")
}
lazyCache, err := utils.NewFnCache(utils.FnCacheConfig{
Context: ctx,
TTL: 15 * time.Second,
Clock: cfg.Clock,
})
if err != nil {
return nil, trace.Wrap(err)
}
m := &UnifiedResourceCache{
logger: slog.With(teleport.ComponentKey, cfg.Component),
cfg: cfg,
nameTree: btree.NewG(cfg.BTreeDegree, func(a, b *item) bool {
return a.Less(b)
}),
typeTree: btree.NewG(cfg.BTreeDegree, func(a, b *item) bool {
return a.Less(b)
}),
resources: make(map[string]resourceCollection),
initializationC: make(chan struct{}),
ResourceGetter: cfg.ResourceGetter,
cache: lazyCache,
stale: true,
}
if err := newWatcher(ctx, m, cfg.ResourceWatcherConfig); err != nil {
return nil, trace.Wrap(err, "creating unified resource watcher")
}
return m, nil
}
// CheckAndSetDefaults checks and sets default values
func (cfg *UnifiedResourceCacheConfig) CheckAndSetDefaults() error {
if cfg.BTreeDegree <= 0 {
cfg.BTreeDegree = 8
}
if cfg.Clock == nil {
cfg.Clock = clockwork.NewRealClock()
}
if cfg.Component == "" {
cfg.Component = teleport.ComponentUnifiedResource
}
return nil
}
func (c *UnifiedResourceCache) putLocked(resource resource) {
key := resourceKey(resource)
sortKey := makeResourceSortKey(resource)
if collection, exists := c.resources[key]; exists {
// If the resource has changed in such a way that the sort keys
// for the nameTree or typeTree change, remove the old entries
// from those trees before adding a new one. This can happen
// when a node's hostname changes
oldSortKey := makeResourceSortKey(collection.get())
if oldSortKey.byName.Compare(sortKey.byName) != 0 {
c.deleteSortKey(oldSortKey)
}
collection.put(resource)
} else {
c.resources[key] = newResourceCollection(resource)
}
c.nameTree.ReplaceOrInsert(&item{Key: sortKey.byName, Value: key})
c.typeTree.ReplaceOrInsert(&item{Key: sortKey.byType, Value: key})
}
func putResources[T resource](cache *UnifiedResourceCache, resources []T) {
for _, resource := range resources {
cache.putLocked(resource)
}
}
func (c *UnifiedResourceCache) deleteSortKey(sortKey resourceSortKey) error {
if _, ok := c.nameTree.Delete(&item{Key: sortKey.byName}); !ok {
return trace.NotFound("key %q is not found in unified cache name sort tree", sortKey.byName.String())
}
if _, ok := c.typeTree.Delete(&item{Key: sortKey.byType}); !ok {
return trace.NotFound("key %q is not found in unified cache type sort tree", sortKey.byType.String())
}
return nil
}
func (c *UnifiedResourceCache) deleteLocked(res types.Resource) error {
key := resourceKey(res)
collection, exists := c.resources[key]
if !exists {
return trace.NotFound("cannot delete resource: key %s not found in unified resource cache", key)
}
if empty := collection.remove(res); empty {
sortKey := makeResourceSortKey(collection.get())
c.deleteSortKey(sortKey)
delete(c.resources, key)
}
return nil
}
func (c *UnifiedResourceCache) getSortTree(sortField string) (*btree.BTreeG[*item], error) {
switch sortField {
case "", sortByName:
return c.nameTree, nil
case sortByKind:
return c.typeTree, nil
default:
return nil, trace.NotImplemented("sorting by %v is not supported in unified resources", sortField)
}
}
type iteratedItem struct {
resource resource
key backend.Key
}
// iterateItems is a helper for iterating the correct cache, in the correct order
// for only the specified kinds. All external iteration APIs are built upon this
// method.
func (c *UnifiedResourceCache) iterateItems(ctx context.Context, start string, sortBy types.SortBy, kinds ...string) iter.Seq2[iteratedItem, error] {
return func(yield func(iteratedItem, error) bool) {
kindsMap := make(map[string]struct{})
for _, k := range kinds {
kindsMap[k] = struct{}{}
}
var startKey backend.Key
if start != "" {
startKey = backend.KeyFromString(start)
}
itemIter := (*btree.BTreeG[*item]).AscendGreaterOrEqual
if sortBy.IsDesc {
itemIter = (*btree.BTreeG[*item]).DescendLessOrEqual
}
var excludedStart bool
const defaultPageSize = 100
items := make([]iteratedItem, 0, defaultPageSize)
for {
items = items[:0]
err := c.read(ctx, func(cache *UnifiedResourceCache) error {
tree, err := cache.getSortTree(sortBy.Field)
if err != nil {
return trace.Wrap(err, "getting sort tree")
}
if startKey.IsZero() {
max, ok := tree.Max()
if sortBy.IsDesc && ok {
startKey = max.Key
} else {
startKey = backend.NewKey("")
}
}
itemIter(tree, &item{Key: startKey}, func(item *item) bool {
if excludedStart {
excludedStart = false
if item.Key.Compare(startKey) <= 0 {
return true
}
}
collection, ok := cache.resources[item.Value]
if !ok {
return true
}
if len(kinds) == 0 || c.itemKindMatches(collection.get(), kindsMap) {
items = append(items, iteratedItem{key: item.Key, resource: collection.get()})
}
if len(items) >= defaultPageSize {
startKey = item.Key
excludedStart = true
return false
}
return true
})
return nil
})
if err != nil {
yield(iteratedItem{}, err)
return
}
for _, i := range items {
if !yield(i, nil) {
return
}
}
if len(items) < defaultPageSize {
return
}
}
}
}
// Resources iterates over all resources from the start key that match
// one of the provided kinds. If no kinds are provided, resources of all supported
// kinds are returned.
func (c *UnifiedResourceCache) Resources(ctx context.Context, start string, sortBy types.SortBy, kinds ...string) iter.Seq2[types.ResourceWithLabels, error] {
return func(yield func(types.ResourceWithLabels, error) bool) {
for item, err := range c.iterateItems(ctx, start, sortBy, kinds...) {
if err != nil {
yield(nil, err)
return
}
if !yield(item.resource.CloneResource(), nil) {
return
}
}
}
}
// UnifiedResourcesIterateParams are parameters that are provided to
// UnifiedResourceCache iterators to alter the iteration behavior.
type UnifiedResourcesIterateParams struct {
Start string
Descending bool
}
// Nodes iterates over all cached nodes starting from the provided key.
func (c *UnifiedResourceCache) Nodes(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.Server, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindNode, types.Server.DeepCopy)
}
// AppServers iterates over all cached app servers starting from the provided key.
func (c *UnifiedResourceCache) AppServers(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.AppServer, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindAppServer, types.AppServer.Copy)
}
// DatabaseServers iterates over all cached database servers starting from the provided key.
func (c *UnifiedResourceCache) DatabaseServers(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.DatabaseServer, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindDatabaseServer, types.DatabaseServer.Copy)
}
// KubernetesServers iterates over all cached Kubernetes servers starting from the provided key.
func (c *UnifiedResourceCache) KubernetesServers(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.KubeServer, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindKubeServer, types.KubeServer.Copy)
}
// WindowsDesktops iterates over all cached windows desktops starting from the provided key.
func (c *UnifiedResourceCache) WindowsDesktops(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.WindowsDesktop, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindWindowsDesktop, types.WindowsDesktop.Copy)
}
// GitServers iterates over all cached git servers starting from the provided key.
func (c *UnifiedResourceCache) GitServers(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.Server, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindGitServer, types.Server.DeepCopy)
}
// SAMLIdPServiceProviders iterates over all cached sAML IdP service providers starting from the provided key.
func (c *UnifiedResourceCache) SAMLIdPServiceProviders(ctx context.Context, params UnifiedResourcesIterateParams) iter.Seq2[types.SAMLIdPServiceProvider, error] {
return iterateUnifiedResourceCache(ctx, c, params, types.KindSAMLIdPServiceProvider, types.SAMLIdPServiceProvider.Copy)
}
func iterateUnifiedResourceCache[T resource](ctx context.Context, c *UnifiedResourceCache, params UnifiedResourcesIterateParams, kind string, cloneFn func(T) T) iter.Seq2[T, error] {
return func(yield func(T, error) bool) {
sortBy := types.SortBy{IsDesc: params.Descending, Field: SortByName}
for i, err := range c.iterateItems(ctx, params.Start, sortBy, kind) {
if err != nil {
var t T
yield(t, err)
return
}
switch r := i.resource.(type) {
case T:
if !yield(cloneFn(r), nil) {
return
}
default:
var t T
yield(t, trace.BadParameter("expected type %T, got %T", t, r))
return
}
}
}
}
// IterateUnifiedResources allows building a custom page of resources. All items within the
// range and limit of the request are passed to the matchFn. Only those resource which
// have a true value returned from the matchFn are included in the returned page.
func (c *UnifiedResourceCache) IterateUnifiedResources(ctx context.Context, matchFn func(types.ResourceWithLabels) (bool, error), req *proto.ListUnifiedResourcesRequest) ([]types.ResourceWithLabels, string, error) {
var resources []types.ResourceWithLabels
for item, err := range c.iterateItems(ctx, req.StartKey, req.SortBy, req.Kinds...) {
if err != nil {
return nil, "", trace.Wrap(err)
}
match, err := matchFn(item.resource)
if err != nil {
return nil, "", trace.Wrap(err)
}
if match {
if req.Limit != backend.NoLimit && len(resources) == int(req.Limit) {
return resources, item.key.String(), nil
}
resources = append(resources, item.resource.CloneResource())
}
}
return resources, "", nil
}
func (c *UnifiedResourceCache) itemKindMatches(r resource, kinds map[string]struct{}) bool {
switch r.GetKind() {
case types.KindNode,
types.KindWindowsDesktop,
types.KindLinuxDesktop,
types.KindGitServer,
types.KindDatabase,
types.KindKubernetesCluster:
_, ok := kinds[r.GetKind()]
return ok
case types.KindIdentityCenterAccount:
if _, ok := kinds[types.KindApp]; ok {
return ok
}
_, ok := kinds[types.KindIdentityCenterAccount]
return ok
case types.KindApp:
if _, ok := kinds[types.KindApp]; ok {
return ok
}
if _, ok := kinds[types.KindAppServer]; ok {
return ok
}
if _, ok := kinds[types.KindMCP]; ok && r.GetSubKind() == types.KindMCP {
return true
}
// TODO(gabrielcorado): support LLM subkind.
_, ok := kinds[types.KindIdentityCenterAccount]
return ok
case types.KindKubeServer:
if _, ok := kinds[types.KindKubernetesCluster]; ok {
return ok
}
_, ok := kinds[types.KindKubeServer]
return ok
case types.KindDatabaseServer:
if _, ok := kinds[types.KindDatabase]; ok {
return ok
}
_, ok := kinds[types.KindDatabaseServer]
return ok
case types.KindSAMLIdPServiceProvider:
_, ok := kinds[types.KindSAMLIdPServiceProvider]
return ok
case types.KindAppServer:
if r.GetSubKind() == types.KindIdentityCenterAccount {
if _, ok := kinds[types.KindIdentityCenterAccount]; ok {
return ok
}
}
if _, ok := kinds[types.KindApp]; ok {
return ok
}
if _, ok := kinds[types.KindAppServer]; ok {
return ok
}
if _, ok := kinds[types.KindMCP]; ok {
type appGetter interface {
GetApp() types.Application
}
if appServer, ok := r.(appGetter); ok && appServer.GetApp().GetSubKind() == types.SubKindMCP {
return true
}
}
// TODO(gabrielcorado): support LLM subkind.
return false
default:
return false
}
}
// GetUnifiedResources returns a list of all resources stored in the current unifiedResourceCollector tree in ascending order
func (c *UnifiedResourceCache) GetUnifiedResources(ctx context.Context) ([]types.ResourceWithLabels, error) {
var resources []types.ResourceWithLabels
for resource, err := range c.Resources(ctx, "", types.SortBy{IsDesc: false, Field: sortByName}) {
if err != nil {
return nil, trace.Wrap(err)
}
resources = append(resources, resource)
}
return resources, nil
}
// GetUnifiedResourcesByIDs will take a list of ids and return any items found in the unifiedResourceCache tree by id and that return true from matchFn
func (c *UnifiedResourceCache) GetUnifiedResourcesByIDs(ctx context.Context, ids []string, matchFn func(types.ResourceWithLabels) (bool, error)) ([]types.ResourceWithLabels, error) {
var resources []types.ResourceWithLabels
err := c.read(ctx, func(cache *UnifiedResourceCache) error {
for _, id := range ids {
key := backend.NewKey(prefix, id)
res, found := cache.nameTree.Get(&item{Key: key})
if !found || res == nil {
continue
}
collection, ok := cache.resources[res.Value]
if !ok {
continue
}
resource := collection.get()
match, err := matchFn(resource)
if err != nil {
return trace.Wrap(err)
}
if match {
resources = append(resources, resource.CloneResource())
}
}
return nil
})
if err != nil {
return nil, trace.Wrap(err)
}
return resources, nil
}
// ResourceGetter is an interface that provides a way to fetch all the resources
// that can be stored in the UnifiedResourceCache
type ResourceGetter interface {
NodesGetter
DatabaseServersGetter
AppServersGetter
WindowsDesktopGetter
LinuxDesktopGetter
KubernetesServerGetter
SAMLIdpServiceProviderGetter
IdentityCenterAccountGetter
GitServerGetter
}
// newWatcher starts and returns a new resource watcher for unified resources.
func newWatcher(ctx context.Context, resourceCache *UnifiedResourceCache, cfg ResourceWatcherConfig) error {
if err := cfg.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err, "setting defaults for unified resource watcher config")
}
if _, err := newResourceWatcher(ctx, resourceCache, cfg); err != nil {
return trace.Wrap(err, "creating a new unified resource watcher")
}
return nil
}
// resourceName is a unique name to be used as a key in the resources map
func resourceKey(resource types.Resource) string {
key := resource.GetName() + "/" + resource.GetKind()
if r, ok := resource.(interface{ GetScope() string }); ok {
if scope := r.GetScope(); scope != "" {
key = scopes.QualifiedName{Name: key, Scope: scope}.String()
}
}
return key
}
type resourceSortKey struct {
byName backend.Key
byType backend.Key
}
// resourceSortKey will generate a key to be used in the sort trees
func makeResourceSortKey(resource types.Resource) resourceSortKey {
var name, kind string
// set the kind to the appropriate "contained" type, rather than
// the container type.
switch r := resource.(type) {
case types.Server:
switch r.GetKind() {
case types.KindNode, types.KindGitServer:
name = r.GetHostname() + "/" + r.GetName()
kind = r.GetKind()
}
case types.AppServer:
app := r.GetApp()
if app != nil {
friendlyName := types.FriendlyName(app)
if friendlyName != "" {
sanitizedFriendlyName := strings.ReplaceAll(types.FriendlyName(app), "/", "-")
// FriendlyName is not unique, and multiple apps may have the same friendly name.
// To prevent collisions in the resource cache, we append the app name to the
// friendly name, ensuring uniqueness.
name = sanitizedFriendlyName + "/" + app.GetName()
} else {
name = app.GetName()
}
if scope := r.GetScope(); scope != "" {
name = scopes.QualifiedName{Name: name, Scope: scope}.String()
}
kind = types.KindApp
}
case types.SAMLIdPServiceProvider:
name = r.GetName()
kind = types.KindApp
case types.KubeServer:
cluster := r.GetCluster()
if cluster != nil {
name = r.GetCluster().GetName()
kind = types.KindKubernetesCluster
}
case types.DatabaseServer:
db := r.GetDatabase()
if db != nil {
name = db.GetName()
kind = types.KindDatabase
}
default:
name = resource.GetName()
kind = resource.GetKind()
}
return resourceSortKey{
// names should be stored as lowercase to keep items sorted as
// expected, regardless of case
byName: backend.NewKey(prefix, strings.ToLower(name), kind),
byType: backend.NewKey(prefix, kind, strings.ToLower(name)),
}
}
func (c *UnifiedResourceCache) getResourcesAndUpdateCurrent(ctx context.Context) error {
newNodes, err := c.getNodes(ctx)
if err != nil {
return trace.Wrap(err)
}
newDbs, err := c.getDatabaseServers(ctx)
if err != nil {
return trace.Wrap(err)
}
newKubes, err := c.getKubeServers(ctx)
if err != nil {
return trace.Wrap(err)
}
newApps, err := c.getAppServers(ctx)
if err != nil {
return trace.Wrap(err)
}
newSAMLApps, err := c.getSAMLApps(ctx)
if err != nil {
return trace.Wrap(err)
}
newDesktops, err := c.getDesktops(ctx)
if err != nil {
return trace.Wrap(err)
}
newLinuxDesktops, err := c.getLinuxDesktops(ctx)
if err != nil {
return trace.Wrap(err)
}
newICAccounts, err := c.getIdentityCenterAccounts(ctx)
if err != nil {
return trace.Wrap(err)
}
newGitServers, err := c.getGitServers(ctx)
if err != nil {
return trace.Wrap(err)
}
c.rw.Lock()
defer c.rw.Unlock()
// empty the trees
c.nameTree.Clear(false)
c.typeTree.Clear(false)
// clear the resource map as well
// c.resources = make(map[string]resource)
clear(c.resources)
putResources(c, newNodes)
putResources(c, newDbs)
putResources(c, newApps)
putResources(c, newKubes)
putResources(c, newSAMLApps)
putResources(c, newDesktops)
putResources(c, newLinuxDesktops)
putResources(c, newICAccounts)
putResources(c, newGitServers)
c.stale = false
c.defineCollectorAsInitialized()
return nil
}
// getNodes will get all nodes
func (c *UnifiedResourceCache) getNodes(ctx context.Context) ([]types.Server, error) {
newNodes, err := c.ResourceGetter.GetNodes(ctx, apidefaults.Namespace)
if err != nil {
return nil, trace.Wrap(err, "getting nodes for unified resource watcher")
}
return newNodes, err
}
// getDatabaseServers will get all database servers
func (c *UnifiedResourceCache) getDatabaseServers(ctx context.Context) ([]types.DatabaseServer, error) {
newDbs, err := c.GetDatabaseServers(ctx, apidefaults.Namespace)
if err != nil {
return nil, trace.Wrap(err, "getting database servers for unified resource watcher")
}
return newDbs, nil
}
// getKubeServers will get all kube servers
func (c *UnifiedResourceCache) getKubeServers(ctx context.Context) ([]types.KubeServer, error) {
newKubes, err := c.GetKubernetesServers(ctx)
if err != nil {
return nil, trace.Wrap(err, "getting kube servers for unified resource watcher")
}
return newKubes, nil
}
// getAppServers will get all application servers
func (c *UnifiedResourceCache) getAppServers(ctx context.Context) ([]types.AppServer, error) {
newApps, err := c.GetApplicationServers(ctx, apidefaults.Namespace)
if err != nil {
return nil, trace.Wrap(err, "getting app servers for unified resource watcher")
}
return newApps, nil
}
// getDesktops will get all windows desktops
func (c *UnifiedResourceCache) getDesktops(ctx context.Context) ([]types.WindowsDesktop, error) {
newDesktops, err := c.GetWindowsDesktops(ctx, types.WindowsDesktopFilter{})
if err != nil {
return nil, trace.Wrap(err, "getting desktops for unified resource watcher")
}
return newDesktops, nil
}
// getLinuxDesktops will get all Linux desktops
func (c *UnifiedResourceCache) getLinuxDesktops(ctx context.Context) ([]resource, error) {
var linuxDesktops []resource
for linuxDesktop, err := range clientutils.Resources(ctx, c.ListLinuxDesktops) {
if err != nil {
return nil, trace.Wrap(err)
}
linuxDesktops = append(linuxDesktops, types.ProtoResource153ToLegacy(linuxDesktop))
}
return linuxDesktops, nil
}
// getSAMLApps will get all SAML Idp Service Providers
func (c *UnifiedResourceCache) getSAMLApps(ctx context.Context) ([]types.SAMLIdPServiceProvider, error) {
var newSAMLApps []types.SAMLIdPServiceProvider
startKey := ""
for {
resp, nextKey, err := c.ListSAMLIdPServiceProviders(ctx, apidefaults.DefaultChunkSize, startKey)
if err != nil {
return nil, trace.Wrap(err, "getting SAML apps for unified resource watcher")
}
newSAMLApps = append(newSAMLApps, resp...)
if nextKey == "" {
break
}
startKey = nextKey
}
return newSAMLApps, nil
}
func (c *UnifiedResourceCache) getIdentityCenterAccounts(ctx context.Context) ([]resource, error) {
var accounts []resource
var startKey string
for {
resp, nextKey, err := c.ListIdentityCenterAccounts(ctx, apidefaults.DefaultChunkSize, startKey)
if err != nil {
return nil, trace.Wrap(err, "getting AWS Identity Center accounts for resource watcher")
}
for _, acct := range resp {
accounts = append(accounts, IdentityCenterAccountToAppServer(acct))
}
if nextKey == "" {
break
}
startKey = nextKey
}
return accounts, nil
}
func (c *UnifiedResourceCache) getGitServers(ctx context.Context) (all []types.Server, err error) {
var page []types.Server
nextToken := ""
for {
page, nextToken, err = c.ListGitServers(ctx, apidefaults.DefaultChunkSize, nextToken)
if err != nil {
return nil, trace.Wrap(err, "getting Git servers for unified resource watcher")
}
all = append(all, page...)
if nextToken == "" {
break
}
}
return all, nil
}
// read applies the supplied closure to either the primary tree or the ttl-based fallback tree depending on
// whether or not the cache is currently healthy. locking is handled internally and the passed-in tree should
// not be accessed after the closure completes.
func (c *UnifiedResourceCache) read(ctx context.Context, fn func(cache *UnifiedResourceCache) error) error {
c.rw.RLock()
if !c.stale {
err := fn(c)
c.rw.RUnlock()
return err
}
c.rw.RUnlock()
ttlCache, err := utils.FnCacheGet(ctx, c.cache, "unified_resources", func(ctx context.Context) (*UnifiedResourceCache, error) {
fallbackCache := &UnifiedResourceCache{
cfg: c.cfg,
nameTree: btree.NewG(c.cfg.BTreeDegree, func(a, b *item) bool {
return a.Less(b)
}),
typeTree: btree.NewG(c.cfg.BTreeDegree, func(a, b *item) bool {
return a.Less(b)
}),
resources: make(map[string]resourceCollection),
ResourceGetter: c.ResourceGetter,
initializationC: make(chan struct{}),
}
if err := fallbackCache.getResourcesAndUpdateCurrent(ctx); err != nil {
return nil, trace.Wrap(err)
}
return fallbackCache, nil
})
c.rw.RLock()
if !c.stale {
// primary became healthy while we were waiting
err := fn(c)
c.rw.RUnlock()
return err
}
c.rw.RUnlock()
if err != nil {
// ttl-tree setup failed
return trace.Wrap(err)
}
err = fn(ttlCache)
return err
}
func (c *UnifiedResourceCache) notifyStale() {
c.rw.Lock()
defer c.rw.Unlock()
c.stale = true
}
func (c *UnifiedResourceCache) initializationChan() <-chan struct{} {
return c.initializationC
}
// IsInitialized is used to check that the cache has done its initial
// sync
func (c *UnifiedResourceCache) IsInitialized() bool {
select {
case <-c.initializationC:
return true
default:
return false
}
}
func (c *UnifiedResourceCache) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
c.rw.Lock()
defer c.rw.Unlock()
if c.stale {
return
}
for _, event := range events {
if event.Resource == nil {
c.logger.WarnContext(ctx, "Unexpected event",
"event_type", event.Type,
"resource_kind", event.Resource.GetKind(),
"resource_name", event.Resource.GetName(),
)
continue
}
switch event.Type {
case types.OpDelete:
switch event.Resource.GetKind() {
case types.KindIdentityCenterAccount:
// IdentityCenterAccountToAppServer stores the entry under
// KindAppServer with the IC account's name; rebuild that
// header so the delete matches the cache key.
c.deleteLocked(&types.ResourceHeader{
Kind: types.KindAppServer,
Metadata: types.Metadata{
Name: event.Resource.GetName(),
},
})
default:
c.deleteLocked(event.Resource)
}
case types.OpPut:
switch r := event.Resource.(type) {
case resource:
c.putLocked(r)
case types.Resource153UnwrapperT[*identitycenterv1.Account]:
c.putLocked(IdentityCenterAccountToAppServer(r.UnwrapT()))
default:
c.logger.WarnContext(ctx, "unsupported Resource type", "resource_type", logutils.TypeAttr(r))
}
default:
c.logger.WarnContext(ctx, "unsupported event type", "event_type", event.Type)
continue
}
}
}
// resourceKinds returns a list of resources to be watched.
func (c *UnifiedResourceCache) resourceKinds() []types.WatchKind {
return []types.WatchKind{
{Kind: types.KindNode},
{Kind: types.KindKubeServer},
{Kind: types.KindDatabaseServer},
{Kind: types.KindAppServer},
{Kind: types.KindWindowsDesktop},
{Kind: types.KindLinuxDesktop},
{Kind: types.KindSAMLIdPServiceProvider},
{Kind: types.KindIdentityCenterAccount},
{Kind: types.KindGitServer},
}
}
func (c *UnifiedResourceCache) defineCollectorAsInitialized() {
c.once.Do(func() {
// mark watcher as initialized.
close(c.initializationC)
})
}
// Less is used for Btree operations,
// returns true if item is less than the other one
func (i *item) Less(iother btree.Item) bool {
switch other := iother.(type) {
case *item:
return i.Key.Compare(other.Key) < 0
default:
return false
}
}
type resource interface {
types.ResourceWithLabels
CloneResource() types.ResourceWithLabels
}
type resourceCollection interface {
get() resource
put(r resource)
// remove removes a resource from the collection and returns true if the
// collection itself should be removed.
remove(r types.Resource) bool
}
func newResourceCollection(r resource) resourceCollection {
switch r := r.(type) {
case types.DatabaseServer:
return newServerResourceCollection(r,
func(srv types.DatabaseServer, servers map[string]types.DatabaseServer) types.DatabaseServer {
return &aggregatedDatabase{
DatabaseServer: srv,
status: aggregateHealthStatuses(servers),
}
})
case types.KubeServer:
return newServerResourceCollection(r,
func(srv types.KubeServer, servers map[string]types.KubeServer) types.KubeServer {
return &aggregatedKube{
KubeServer: srv,
status: aggregateHealthStatuses(servers),
}
})
case types.AppServer:
return newServerResourceCollection(r,
func(srv types.AppServer, servers map[string]types.AppServer) types.AppServer {
return &aggregatedAppServer{
AppServer: srv,
features: intersectComponentFeaturesForAppServers(servers),
}
})
case serverResource:
return newServerResourceCollection(r, nil)
default:
return &singularResourceCollection{latest: r}
}
}
func aggregateHealthStatuses[T types.TargetHealthStatusGetter](hgs map[string]T) types.TargetHealthStatus {
return types.AggregateHealthStatus(func(yield func(types.TargetHealthStatus) bool) {
for _, hg := range hgs {
if !yield(hg.GetTargetHealthStatus()) {
return
}
}
})
}
type singularResourceCollection struct {
latest resource
}
func (c *singularResourceCollection) get() resource { return c.latest }
func (c *singularResourceCollection) put(r resource) { c.latest = r }
func (c *singularResourceCollection) remove(types.Resource) bool { return true }
// serverResource is a type of resource that may have multiple agents
// heartbeating it.
type serverResource interface {
resource
GetHostID() string
}
type serverResourceCollection[R serverResource] struct {
aggregate R
aggregationFn func(latest R, servers map[string]R) R
servers map[string]R
}
func newServerResourceCollection[R serverResource](r R, aggFn func(latest R, servers map[string]R) R) *serverResourceCollection[R] {
if aggFn == nil {
aggFn = func(r R, _ map[string]R) R {
return r
}
}
collection := &serverResourceCollection[R]{
servers: make(map[string]R),
aggregationFn: aggFn,
}
collection.put(r)
return collection
}
func (c *serverResourceCollection[R]) get() resource {
return c.aggregate
}
func (c *serverResourceCollection[R]) put(r resource) {
if r, ok := r.(R); ok {
c.servers[r.GetHostID()] = r
c.aggregate = c.aggregationFn(r, c.servers)
}
}
func (c *serverResourceCollection[R]) remove(r types.Resource) bool {
// This looks insane, but we only get a [types.ResourceHeader] in
// [types.OpDelete] events.
// The types that actually implement [resourceServer] all store the host ID
// in the description of the resource header metadata on deletion.
// If a new type is added that implements [resourceServer] and the
// unified resource watchers starts watching it, then please:
// - add it to the isResourceServer helper func
// - ensure host ID is stored in the metadata description in its event parser
// - add test coverage for it in TestUnifiedResourceWatcher_DeleteEvent
delete(c.servers, r.GetMetadata().Description)
for _, s := range c.servers {
c.aggregate = c.aggregationFn(s, c.servers)
return false
}
return true
}
// aggregatedAppServer wraps an app server with aggregated ComponentFeatures
// in order to perform an intersection of all features reported by multiple
// AppServers serving the same app. Only ComponentFeatures supported by *all*
// AppServers will be reported in the aggregated Resource.
type aggregatedAppServer struct {
types.AppServer
features *componentfeaturesv1.ComponentFeatures
}
// Copy returns a copy of the underlying app server with aggregated
// [componentfeaturesv1.ComponentFeatures].
func (a *aggregatedAppServer) Copy() types.AppServer {
out := a.AppServer.Copy()
out.SetComponentFeatures(a.GetComponentFeatures())
return out
}
// CloneResource returns a copy of the underlying app server with
// aggregated [componentfeaturesv1.ComponentFeatures].
func (a *aggregatedAppServer) CloneResource() types.ResourceWithLabels {
return a.Copy()
}
func (a *aggregatedAppServer) GetComponentFeatures() *componentfeaturesv1.ComponentFeatures {
if a.features == nil {
return nil
}
return componentfeatures.Join(a.features)
}
func intersectComponentFeaturesForAppServers(servers map[string]types.AppServer) *componentfeaturesv1.ComponentFeatures {
allFeatures := make([]*componentfeaturesv1.ComponentFeatures, 0, len(servers))
for _, s := range servers {
allFeatures = append(allFeatures, componentfeatures.GetEffectiveServerFeatures(s))
}
return componentfeatures.Intersect(allFeatures...)
}
// aggregatedDatabase wraps a database server with aggregated health status.
// It is assumed that multiple heartbeats with the same resource name but
// different host IDs may be received and they may report different health
// statuses.
// This type provides the following properties:
// - avoid cloning the resource *before* filtering.
// - when the resource is cloned *after* filtering, set the clone's health
// status to the aggregate health status.
//
// Go generics do not support embedding a generic type, otherwise this type
// would be made generic.
type aggregatedDatabase struct {
types.DatabaseServer
status types.TargetHealthStatus
}
// This type MUST implement [types.DatabaseServer] to act as a facade type,
// otherwise dynamic assertions elsewhere will fail.
var _ types.DatabaseServer = (*aggregatedDatabase)(nil)
// GetTargetHealthStatus gets the aggregate health status for filtering by
// health status.
func (d *aggregatedDatabase) GetTargetHealthStatus() types.TargetHealthStatus {
return d.status
}
// Copy returns a copy of the underlying database server with aggregated health
// status.
func (d *aggregatedDatabase) Copy() types.DatabaseServer {
out := d.DatabaseServer.Copy()
out.SetTargetHealthStatus(d.status)
return out
}
// CloneResource returns a copy of the underlying database server with
// aggregated health status.
func (d *aggregatedDatabase) CloneResource() types.ResourceWithLabels {
return d.Copy()
}
// aggregatedKube wraps a kube server with aggregated health status.
// It is assumed that multiple heartbeats with the same resource name but
// different host IDs may be received and they may report different health
// statuses.
// This type provides the following properties:
// - avoid cloning the resource *before* filtering.
// - when the resource is cloned *after* filtering, set the clone's health
// status to the aggregate health status.
//
// Go generics do not support embedding a generic type, otherwise this type
// would be made generic.
type aggregatedKube struct {
types.KubeServer
status types.TargetHealthStatus
}
// This type MUST implement [types.KubeServer] to act as a facade type,
// otherwise dynamic assertions elsewhere will fail.
var _ types.KubeServer = (*aggregatedKube)(nil)
// GetTargetHealthStatus gets the aggregate health status for filtering by
// health status.
func (d *aggregatedKube) GetTargetHealthStatus() types.TargetHealthStatus {
return d.status
}
// Copy returns a copy of the underlying kube server with aggregated health
// status.
func (d *aggregatedKube) Copy() types.KubeServer {
out := d.KubeServer.Copy()
out.SetTargetHealthStatus(d.status)
return out
}
// CloneResource returns a copy of the underlying kube server with
// aggregated health status.
func (d *aggregatedKube) CloneResource() types.ResourceWithLabels {
return d.Copy()
}
type item struct {
// Key is a key of the key value item. This will be different based on which sorting tree
// the item is in
Key backend.Key
// Value will be the resourceKey used in the resources map to get the resource
Value string
}
const (
prefix = "unified_resource"
sortByName string = "name"
sortByKind string = "kind"
)
// MakePaginatedResource converts a resource into a paginated proto representation.
func MakePaginatedResource(requestType string, r types.ResourceWithLabels, requiresRequest bool) (*proto.PaginatedResource, error) {
var protoResource *proto.PaginatedResource
resourceKind := requestType
if requestType == types.KindUnifiedResource {
resourceKind = r.GetKind()
}
var logins []string
resource := r
if enriched, ok := r.(*types.EnrichedResource); ok {
resource = enriched.ResourceWithLabels
logins = enriched.Logins
}
switch resourceKind {
case types.KindDatabaseServer:
database, ok := resource.(*types.DatabaseServerV3)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_DatabaseServer{DatabaseServer: database}, RequiresRequest: requiresRequest}
case types.KindDatabaseService:
databaseService, ok := resource.(*types.DatabaseServiceV1)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_DatabaseService{DatabaseService: databaseService}, RequiresRequest: requiresRequest}
case types.KindAppServer:
app, ok := resource.(*types.AppServerV3)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_AppServer{AppServer: app}, Logins: logins, RequiresRequest: requiresRequest}
case types.KindNode:
srv, ok := resource.(*types.ServerV2)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_Node{Node: srv}, Logins: logins, RequiresRequest: requiresRequest}
case types.KindKubeServer:
srv, ok := resource.(*types.KubernetesServerV3)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_KubernetesServer{KubernetesServer: srv}, RequiresRequest: requiresRequest}
case types.KindLinuxDesktop:
unwrapper, ok := resource.(types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop])
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: proto.PackLinuxDesktop(unwrapper.UnwrapT()), Logins: logins, RequiresRequest: requiresRequest}
case types.KindWindowsDesktop:
desktop, ok := resource.(*types.WindowsDesktopV3)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_WindowsDesktop{WindowsDesktop: desktop}, Logins: logins, RequiresRequest: requiresRequest}
case types.KindWindowsDesktopService:
desktopService, ok := resource.(*types.WindowsDesktopServiceV3)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_WindowsDesktopService{WindowsDesktopService: desktopService}, RequiresRequest: requiresRequest}
case types.KindKubernetesCluster:
cluster, ok := resource.(*types.KubernetesClusterV3)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_KubeCluster{KubeCluster: cluster}, RequiresRequest: requiresRequest}
case types.KindUserGroup:
userGroup, ok := resource.(*types.UserGroupV1)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{Resource: &proto.PaginatedResource_UserGroup{UserGroup: userGroup}, RequiresRequest: requiresRequest}
case types.KindSAMLIdPServiceProvider:
serviceProvider, ok := resource.(*types.SAMLIdPServiceProviderV1)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{
Resource: &proto.PaginatedResource_SAMLIdPServiceProvider{
SAMLIdPServiceProvider: serviceProvider,
},
RequiresRequest: requiresRequest,
}
case types.KindIdentityCenterAccount:
unwrapper, ok := resource.(types.Resource153UnwrapperT[IdentityCenterAccount])
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{
Resource: &proto.PaginatedResource_AppServer{
AppServer: IdentityCenterAccountToAppServer(unwrapper.UnwrapT().Account),
},
RequiresRequest: requiresRequest,
}
case types.KindIdentityCenterAccountAssignment:
unwrapper, ok := resource.(types.Resource153UnwrapperT[IdentityCenterAccountAssignment])
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{
Resource: proto.PackICAccountAssignment(unwrapper.UnwrapT().AccountAssignment),
RequiresRequest: requiresRequest,
}
case types.KindGitServer:
server, ok := resource.(*types.ServerV2)
if !ok {
return nil, trace.BadParameter("%s has invalid type %T", resourceKind, resource)
}
protoResource = &proto.PaginatedResource{
Resource: &proto.PaginatedResource_GitServer{
GitServer: server,
},
RequiresRequest: requiresRequest,
}
default:
return nil, trace.NotImplemented("resource type %s doesn't support pagination", resource.GetKind())
}
return protoResource, nil
}
// MakePaginatedResources converts a list of resources into a list of paginated proto representations.
func MakePaginatedResources(requestType string, resources []types.ResourceWithLabels, requestableMap map[string]struct{}) ([]*proto.PaginatedResource, error) {
paginatedResources := make([]*proto.PaginatedResource, 0, len(resources))
for _, r := range resources {
_, requiresRequest := requestableMap[r.GetName()]
protoResource, err := MakePaginatedResource(requestType, r, requiresRequest)
if err != nil {
return nil, trace.Wrap(err)
}
paginatedResources = append(paginatedResources, protoResource)
}
return paginatedResources, nil
}
const (
SortByName string = "name"
SortByKind string = "kind"
)
/*
* 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"
"encoding/json"
"fmt"
"strings"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// ValidateUser validates the User and sets default values
func ValidateUser(u types.User) error {
if err := CheckAndSetDefaults(u); err != nil {
return trace.Wrap(err)
}
if localAuth := u.GetLocalAuth(); localAuth != nil {
if err := ValidateLocalAuthSecrets(localAuth); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// ValidateUserRoles checks that all the roles in the user exist
func ValidateUserRoles(ctx context.Context, u types.User, roleGetter RoleGetter) error {
for _, role := range u.GetRoles() {
if _, err := roleGetter.GetRole(ctx, role); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// LoginAttempt represents successful or unsuccessful attempt for user to login
type LoginAttempt struct {
// Time is time of the attempt
Time time.Time `json:"time"`
// Success indicates whether attempt was successful
Success bool `json:"bool"`
}
// Check checks parameters
func (la *LoginAttempt) Check() error {
if la.Time.IsZero() {
return trace.BadParameter("missing parameter time")
}
return nil
}
// UnmarshalUser unmarshals the User resource from JSON.
func UnmarshalUser(bytes []byte, opts ...MarshalOption) (*types.UserV2, error) {
var h types.ResourceHeader
err := json.Unmarshal(bytes, &h)
if err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V2:
var u types.UserV2
if err := utils.FastUnmarshal(bytes, &u); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := ValidateUser(&u); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
u.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
u.SetExpiry(cfg.Expires)
}
return &u, nil
}
return nil, trace.BadParameter("user resource version %v is not supported", h.Version)
}
// MarshalUser marshals the User resource to JSON.
func MarshalUser(user types.User, opts ...MarshalOption) ([]byte, error) {
if err := ValidateUser(user); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch user := user.(type) {
case *types.UserV2:
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, user))
default:
return nil, trace.BadParameter("unrecognized user version %T", user)
}
}
// UsernameForRemoteCluster returns an username that is prefixed with "remote-"
// and suffixed with cluster name with the hope that it does not match a real
// local user.
func UsernameForRemoteCluster(localUsername, localClusterName string) string {
return fmt.Sprintf("remote-%v-%v", localUsername, localClusterName)
}
// UsernameForClusterConfig is a configuration struct for UsernameForCluster.
type UsernameForClusterConfig struct {
// User is the username.
User string
// OriginClusterName is the cluster name where the user is authenticated.
OriginClusterName string
// LocalClusterName is the local cluster name.
LocalClusterName string
}
// UsernameForCluster returns an username that is prefixed with "remote-"
// and suffixed with cluster name if the user is from a remote cluster,
// otherwise returns the local username.
func UsernameForCluster(cfg UsernameForClusterConfig) string {
// originClusterName == "" is a special case for backward compatibility
// with older clients that do not send origin cluster name.
// In this case we assume the user is local.
if cfg.OriginClusterName == cfg.LocalClusterName || cfg.OriginClusterName == "" {
return cfg.User
}
return UsernameForRemoteCluster(cfg.User, cfg.OriginClusterName)
}
// ResolveUserDisplays resolves usernames to display values keyed by username,
// reading each unique name once through getter via types.User.GetDisplay.
//
// Missing users are absent from the result. A user with no distinct display is
// present with a zero-value types.UserDisplay. Blank usernames are skipped, and
// any non-NotFound error aborts with no partial map.
func ResolveUserDisplays(ctx context.Context, getter UserGetter, usernames []string) (map[string]types.UserDisplay, error) {
displays := make(map[string]types.UserDisplay)
seen := make(map[string]struct{})
for _, username := range usernames {
if strings.TrimSpace(username) == "" {
continue // skipping whitespace-only username
}
if _, ok := seen[username]; ok {
continue // skipping duplicate username
}
seen[username] = struct{}{}
user, err := getter.GetUser(ctx, username, false)
if trace.IsNotFound(err) {
continue // skipping missing user
}
if err != nil {
return nil, trace.Wrap(err, "resolving display for user %q", username)
}
displays[username] = user.GetDisplay()
}
return displays, nil
}
/*
* 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"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types/userloginstate"
"github.com/gravitational/teleport/lib/utils"
)
// UserLoginStatesGetter is the interface for reading user login states.
type UserLoginStatesGetter interface {
// GetUserLoginStates returns the all user login state resources.
GetUserLoginStates(context.Context) ([]*userloginstate.UserLoginState, error)
// GetUserLoginState returns the specified user login state resource.
GetUserLoginState(context.Context, string) (*userloginstate.UserLoginState, error)
// ListUserLoginStates returns a paginated list of user login state resources.
ListUserLoginStates(ctx context.Context, pageSize int, nextToken string) ([]*userloginstate.UserLoginState, string, error)
}
// UserLoginStates is the interface for managing with user login states.
type UserLoginStates interface {
UserLoginStatesGetter
// UpsertUserLoginState creates or updates a user login state resource.
UpsertUserLoginState(context.Context, *userloginstate.UserLoginState) (*userloginstate.UserLoginState, error)
// DeleteUserLoginState removes the specified user login state resource.
DeleteUserLoginState(context.Context, string) error
// DeleteAllUserLoginStates removes all user login state resources.
DeleteAllUserLoginStates(context.Context) error
}
// MarshalUserLoginState marshals the user login state resource to JSON.
func MarshalUserLoginState(userLoginState *userloginstate.UserLoginState, opts ...MarshalOption) ([]byte, error) {
if err := userLoginState.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
prevRev := userLoginState.GetRevision()
defer func() {
userLoginState.SetRevision(prevRev)
}()
userLoginState.SetRevision("")
}
return utils.FastMarshal(userLoginState)
}
// UnmarshalUserLoginState unmarshals the user login state resource from JSON.
func UnmarshalUserLoginState(data []byte, opts ...MarshalOption) (*userloginstate.UserLoginState, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
uls := &userloginstate.UserLoginState{}
if err := utils.FastUnmarshal(data, &uls); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := uls.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
uls.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
uls.SetExpiry(cfg.Expires)
}
return uls, nil
}
// UserOrLoginStateGetter defines an interface that can get user login states or users.
type UserOrLoginStateGetter interface {
UserLoginStatesGetter
UserGetter
}
// GetUserOrLoginState will return the given user login state or if not found, the user itself.
func GetUserOrLoginState(ctx context.Context, getter UserOrLoginStateGetter, username string) (UserState, error) {
uls, err := getter.GetUserLoginState(ctx, username)
if err != nil && !trace.IsNotFound(err) && !trace.IsAccessDenied(err) {
return nil, trace.Wrap(err)
}
if err == nil {
if strings.HasPrefix(username, BotUserPrefix) {
// Bots should never have ULS, but bugs elsewhere can mistakenly
// introduce it. If ULS happens to exist for a user with the bot
// label, we need to ignore it and return the user instead.
// Unfortunately, that means we can't trust the labels of the ULS we
// just fetched - they might be invalid/empty and missing the bot
// label.
// However, if the user has the `bot-` prefix, we know the user is
// plausibly a bot, so we can fetch the user early and check
// directly.
botUser, botErr := getter.GetUser(ctx, username, false)
if botErr == nil && botUser.IsBot() {
// User was fetchable and has the bot label, return it
// immediately.
return botUser, nil
}
// If the user was not fetchable or wasn't a bot, fall back to the
// standard logic.
}
return uls, nil
}
user, err := getter.GetUser(ctx, username, false)
return user, trace.Wrap(err)
}
/*
* Teleport
* Copyright (C) 2026 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"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
apidefaults "github.com/gravitational/teleport/api/defaults"
userspb "github.com/gravitational/teleport/api/gen/proto/go/teleport/users/v1"
"github.com/gravitational/teleport/api/types"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/set"
)
var userSearchLogger = logutils.NewPackageLogger(teleport.ComponentKey, teleport.Component("user", "search"))
// UserSearchLister lists users for display-value search.
type UserSearchLister interface {
ListUsers(ctx context.Context, req *userspb.ListUsersRequest) (*userspb.ListUsersResponse, error)
}
// findUsernamesBySearchKeywords returns usernames whose resolved display values match the keywords.
func findUsernamesBySearchKeywords(ctx context.Context, users UserSearchLister, searchKeywords []string) (set.Set[string], error) {
if len(searchKeywords) == 0 {
return nil, nil
}
usernames := set.NewWithCapacity[string](apidefaults.DefaultChunkSize)
var pageToken string
for {
rsp, err := users.ListUsers(ctx, userspb.ListUsersRequest_builder{
PageSize: apidefaults.DefaultChunkSize,
PageToken: pageToken,
Filter: &types.UserFilter{
SearchKeywords: searchKeywords,
SkipSystemUsers: true,
},
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
for _, user := range rsp.GetUsers() {
display := user.GetDisplay()
// Exclude non-display trait matches.
if !types.MatchSearch([]string{display.Primary, display.Secondary}, searchKeywords, nil) {
continue
}
usernames.Add(user.GetName())
if usernames.Len() == apidefaults.DefaultChunkSize {
// Cap resolved usernames to avoid excessive paging.
return usernames, nil
}
}
pageToken = rsp.GetNextPageToken()
if pageToken == "" {
return usernames, nil
}
}
}
type searchKeywordUsernameResolver struct {
users UserSearchLister
// usernamesBySearchKeyword caches the resolved usernames for each keyword.
usernamesBySearchKeyword map[string]set.Set[string]
}
// NewSearchKeywordUsernameResolver returns a memoizing resolver for search-keyword username matches.
func NewSearchKeywordUsernameResolver(users UserSearchLister) func(context.Context, string) set.Set[string] {
resolver := &searchKeywordUsernameResolver{
users: users,
usernamesBySearchKeyword: make(map[string]set.Set[string]),
}
return resolver.resolveUsernames
}
func (r *searchKeywordUsernameResolver) resolveUsernames(ctx context.Context, searchKeyword string) set.Set[string] {
searchKeyword = strings.TrimSpace(searchKeyword)
if searchKeyword == "" {
return nil
}
if usernames, ok := r.usernamesBySearchKeyword[searchKeyword]; ok {
return usernames
}
usernames, err := findUsernamesBySearchKeywords(ctx, r.users, []string{searchKeyword})
if err != nil {
userSearchLogger.WarnContext(ctx, "Failed to resolve search keyword to users",
"search_keywords", []string{searchKeyword},
"error", err,
)
r.usernamesBySearchKeyword[searchKeyword] = nil
return nil
}
r.usernamesBySearchKeyword[searchKeyword] = usernames
return usernames
}
// NewAccessRequestSearchMatcher returns a matcher that checks stored request fields and requester user-search matches.
func NewAccessRequestSearchMatcher(searchKeywords []string, users UserSearchLister) func(context.Context, *types.AccessRequestV3) bool {
resolveToUsernames := NewSearchKeywordUsernameResolver(users)
return func(ctx context.Context, accessRequest *types.AccessRequestV3) bool {
return types.MatchSearch(accessRequest.SearchableFields(), searchKeywords, func(searchKeyword string) bool {
return resolveToUsernames(ctx, searchKeyword).Contains(accessRequest.GetUser())
})
}
}
/*
* Teleport
* Copyright (C) 2024 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"
usertasksv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/usertasks/v1"
)
// UserTasks is the interface for managing user tasks resources.
type UserTasks interface {
// CreateUserTask creates a new user tasks resource.
CreateUserTask(context.Context, *usertasksv1.UserTask) (*usertasksv1.UserTask, error)
// UpsertUserTask creates or updates the user tasks resource.
UpsertUserTask(context.Context, *usertasksv1.UserTask) (*usertasksv1.UserTask, error)
// GetUserTask returns the user tasks resource by name.
GetUserTask(ctx context.Context, name string) (*usertasksv1.UserTask, error)
// ListUserTasks returns the user tasks resources.
ListUserTasks(ctx context.Context, pageSize int64, nextToken string, filters *usertasksv1.ListUserTasksFilters) ([]*usertasksv1.UserTask, string, error)
// UpdateUserTask updates the user tasks resource.
UpdateUserTask(context.Context, *usertasksv1.UserTask) (*usertasksv1.UserTask, error)
// DeleteUserTask deletes the user tasks resource by name.
DeleteUserTask(context.Context, string) error
}
// MarshalUserTask marshals the UserTask object into a JSON byte array.
func MarshalUserTask(object *usertasksv1.UserTask, opts ...MarshalOption) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalUserTask unmarshals the UserTask object from a JSON byte array.
func UnmarshalUserTask(data []byte, opts ...MarshalOption) (*usertasksv1.UserTask, error) {
return UnmarshalProtoResource[*usertasksv1.UserTask](data, opts...)
}
func MatchUserTask(ut *usertasksv1.UserTask, filters *usertasksv1.ListUserTasksFilters) bool {
integrationFilter := filters.GetIntegration()
if integrationFilter != "" && integrationFilter != ut.GetSpec().GetIntegration() {
return false
}
stateFilter := filters.GetTaskState()
if stateFilter != "" && stateFilter != ut.GetSpec().GetState() {
return false
}
return true
}
/*
* 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 (
"github.com/gravitational/teleport/api/client/proto"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/entitlements"
"github.com/gravitational/teleport/lib/modules"
)
type ResourceAccess struct {
List bool `json:"list"`
Read bool `json:"read"`
Edit bool `json:"edit"`
Create bool `json:"create"`
Delete bool `json:"remove"`
Use bool `json:"use"`
}
// MobileDeviceAccess defines permissions for the mobile_device resource.
// It uses a dedicated shape rather than ResourceAccess because mobile_device
// exposes a custom verb, not the standard list/read/edit/create/delete/use set.
type MobileDeviceAccess struct {
// CreateEnrollToken reflects the mobile_device.create_enroll_token verb,
// which gates a user's ability to start mobile device enrollment.
CreateEnrollToken bool `json:"createEnrollToken"`
}
// UserACL is derived from a user's role set and includes
// information as to what features the user is allowed to use.
type UserACL struct {
// RecordedSessions defines access to recorded sessions.
RecordedSessions ResourceAccess `json:"recordedSessions"`
// ActiveSessions defines access to active sessions.
ActiveSessions ResourceAccess `json:"activeSessions"`
// AuthConnectors defines access to auth.connectors.
AuthConnectors ResourceAccess `json:"authConnectors"`
// Roles defines access to roles.
Roles ResourceAccess `json:"roles"`
// Users defines access to users.
Users ResourceAccess `json:"users"`
// TrustedClusters defines access to trusted clusters.
TrustedClusters ResourceAccess `json:"trustedClusters"`
// Events defines access to audit logs.
Events ResourceAccess `json:"events"`
// Tokens defines access to tokens.
Tokens ResourceAccess `json:"tokens"`
// Nodes defines access to nodes.
Nodes ResourceAccess `json:"nodes"`
// AppServers defines access to application servers
AppServers ResourceAccess `json:"appServers"`
// DBServers defines access to database servers.
DBServers ResourceAccess `json:"dbServers"`
// DB defines access to database resource.
DB ResourceAccess `json:"db"`
// KubeServers defines access to kubernetes servers.
KubeServers ResourceAccess `json:"kubeServers"`
// Desktops defines access to desktops.
Desktops ResourceAccess `json:"desktops"`
// AccessRequests defines access to access requests.
AccessRequests ResourceAccess `json:"accessRequests"`
// Billing defines access to billing information.
Billing ResourceAccess `json:"billing"`
// ConnectionDiagnostic defines access to connection diagnostics.
ConnectionDiagnostic ResourceAccess `json:"connectionDiagnostic"`
// Clipboard defines whether the user can use a shared clipboard during windows desktop sessions.
Clipboard bool `json:"clipboard"`
// DesktopSessionRecording defines whether the user's desktop sessions are being recorded.
DesktopSessionRecording bool `json:"desktopSessionRecording"`
// DirectorySharing defines whether a user is permitted to share a directory during windows desktop sessions.
DirectorySharing bool `json:"directorySharing"`
// Download defines whether the user has access to download Teleport Enterprise Binaries
Download ResourceAccess `json:"download"`
// Download defines whether the user has access to download the license
License ResourceAccess `json:"license"`
// Plugins defines whether the user has access to manage hosted plugin instances
Plugins ResourceAccess `json:"plugins"`
// Integrations defines whether the user has access to manage integrations.
Integrations ResourceAccess `json:"integrations"`
// UserTasks defines whether the user has access to manage UserTasks.
UserTasks ResourceAccess `json:"userTasks"`
// DeviceTrust defines access to device trust.
DeviceTrust ResourceAccess `json:"deviceTrust"`
// Locks defines access to locking resources.
Locks ResourceAccess `json:"lock"`
// SAMLIdpServiceProvider defines access to `saml_idp_service_provider` objects.
SAMLIdpServiceProvider ResourceAccess `json:"samlIdpServiceProvider"`
// AccessList defines access to access list management.
AccessList ResourceAccess `json:"accessList"`
// DiscoveryConfig defines whether the user has access to manage DiscoveryConfigs.
DiscoveryConfig ResourceAccess `json:"discoverConfigs"`
// AuditQuery defines access to audit query management.
AuditQuery ResourceAccess `json:"auditQuery"`
// SecurityReport defines access to security reports.
SecurityReport ResourceAccess `json:"securityReport"`
// ExternalAuditStorage defines access to manage ExternalAuditStorage
ExternalAuditStorage ResourceAccess `json:"externalAuditStorage"`
// AccessGraph defines access to access graph.
AccessGraph ResourceAccess `json:"accessGraph"`
// Bots defines access to manage Bots.
Bots ResourceAccess `json:"bots"`
// BotInstances defines access to manage bot instances
BotInstances ResourceAccess `json:"botInstances"`
// Instances defines access to manage instances
Instances ResourceAccess `json:"instances"`
// AccessMonitoringRule defines access to manage access monitoring rule resources.
AccessMonitoringRule ResourceAccess `json:"accessMonitoringRule"`
// CrownJewel defines access to manage CrownJewel resources.
CrownJewel ResourceAccess `json:"crownJewel"`
// AccessGraphSettings defines access to manage access graph settings.
AccessGraphSettings ResourceAccess `json:"accessGraphSettings"`
// ReviewRequests defines the ability to review requests
ReviewRequests bool `json:"reviewRequests"`
// Contact defines the ability to manage contacts
Contact ResourceAccess `json:"contact"`
// FileTransferAccess defines the ability to perform remote file operations via SCP or SFTP
FileTransferAccess bool `json:"fileTransferAccess"`
// WebTerminalClipboardMode determines clipboard behavior in the Web UI terminal.
WebTerminalClipboardMode types.WebTerminalClipboardMode `json:"webTerminalClipboardMode,omitempty"`
// GitServers defines access to Git servers.
GitServers ResourceAccess `json:"gitServers"`
// WorkloadIdentity defines access to Workload Identity
WorkloadIdentity ResourceAccess `json:"workloadIdentity"`
// ClientIPRestriction defines access to Cloud IP Restrictions
ClientIPRestriction ResourceAccess `json:"clientIpRestriction"`
// InferenceModel defines access to session summaries inference model.
InferenceModel ResourceAccess `json:"inferenceModel"`
// InferencePolicy defines access to session summaries inference policy.
InferencePolicy ResourceAccess `json:"inferencePolicy"`
// InferenceSecret defines access to session summaries inference secret.
InferenceSecret ResourceAccess `json:"inferenceSecret"`
// Classifier defines access to session summarization classifiers.
Classifier ResourceAccess `json:"classifier"`
// AutoUpdateConfig defines access to autoupdate config.
AutoUpdateConfig ResourceAccess `json:"autoUpdateConfig"`
// AutoUpdateVersion defines access to autoupdate version.
AutoUpdateVersion ResourceAccess `json:"autoUpdateVersion"`
// AutoUpdateAgentRollout defines access to autoupdate agent rollout.
AutoUpdateAgentRollout ResourceAccess `json:"autoUpdateAgentRollout"`
// AutoUpdateAgentReport defines access to autoupdate agent reports.
AutoUpdateAgentReport ResourceAccess `json:"autoUpdateAgentReport"`
// Beam defines access to Beams
Beam ResourceAccess `json:"beam"`
// MobileDevice defines permissions for the mobile_device resource.
MobileDevice MobileDeviceAccess `json:"mobileDevice"`
}
func hasAccess(roleSet RoleSet, ctx *Context, kind string, verbs ...string) bool {
for _, verb := range verbs {
// Since this check occurs often and does not imply the caller is trying to
// ResourceAccess any resource, silence any logging done on the proxy.
if err := roleSet.GuessIfAccessIsPossible(ctx, apidefaults.Namespace, kind, verb); err != nil {
return false
}
}
return true
}
func newAccess(roleSet RoleSet, ctx *Context, kind string) ResourceAccess {
return ResourceAccess{
List: hasAccess(roleSet, ctx, kind, types.VerbList),
Read: hasAccess(roleSet, ctx, kind, types.VerbRead),
Edit: hasAccess(roleSet, ctx, kind, types.VerbUpdate),
Create: hasAccess(roleSet, ctx, kind, types.VerbCreate),
Delete: hasAccess(roleSet, ctx, kind, types.VerbDelete),
Use: hasAccess(roleSet, ctx, kind, types.VerbUse),
}
}
// NewUserACL builds an ACL for a user based on their roles.
func NewUserACL(user types.User, userRoles RoleSet, features proto.Features, desktopRecordingEnabled, accessMonitoringEnabled bool) UserACL {
ctx := &Context{User: user}
recordedSessionAccess := newAccess(userRoles, ctx, types.KindSession)
roleAccess := newAccess(userRoles, ctx, types.KindRole)
authConnectors := newAccess(userRoles, ctx, types.KindAuthConnector)
trustedClusterAccess := newAccess(userRoles, ctx, types.KindTrustedCluster)
eventAccess := newAccess(userRoles, ctx, types.KindEvent)
userAccess := newAccess(userRoles, ctx, types.KindUser)
tokenAccess := newAccess(userRoles, ctx, types.KindToken)
nodeAccess := newAccess(userRoles, ctx, types.KindNode)
appServerAccess := newAccess(userRoles, ctx, types.KindAppServer)
dbServerAccess := newAccess(userRoles, ctx, types.KindDatabaseServer)
dbAccess := newAccess(userRoles, ctx, types.KindDatabase)
kubeServerAccess := newAccess(userRoles, ctx, types.KindKubeServer)
requestAccess := newAccess(userRoles, ctx, types.KindAccessRequest)
accessMonitoringRules := newAccess(userRoles, ctx, types.KindAccessMonitoringRule)
desktopAccess := newAccess(userRoles, ctx, types.KindWindowsDesktop)
cnDiagnosticAccess := newAccess(userRoles, ctx, types.KindConnectionDiagnostic)
samlIdpServiceProviderAccess := newAccess(userRoles, ctx, types.KindSAMLIdPServiceProvider)
gitServersAccess := newAccess(userRoles, ctx, types.KindGitServer)
// active sessions are a special case - if a user's role set has any join_sessions
// policies then the ACL must permit showing active sessions
activeSessionAccess := newAccess(userRoles, ctx, types.KindSSHSession)
if userRoles.CanJoinSessions() {
activeSessionAccess.List = true
activeSessionAccess.Read = true
}
// The billing dashboards are available in: cloud clusters &
// usage-based self-hosted non-stripe dashboards.
var billingAccess ResourceAccess
isDashboard := IsDashboard(features)
isUsageBased := features.IsUsageBased
isStripeManaged := features.IsStripeManaged
if features.Cloud || (isDashboard && isUsageBased && !isStripeManaged) {
billingAccess = newAccess(userRoles, ctx, types.KindBilling)
}
var pluginsAccess ResourceAccess
if features.Plugins {
pluginsAccess = newAccess(userRoles, ctx, types.KindPlugin)
}
var accessGraphAccess ResourceAccess
var accessGraphSettings ResourceAccess
if features.AccessGraph {
accessGraphAccess = newAccess(userRoles, ctx, types.KindAccessGraph)
}
// accessGraphSettings should always be enabled for users to interact with demo mode
accessGraphSettings = newAccess(userRoles, ctx, types.KindAccessGraphSettings)
clipboard := userRoles.DesktopClipboard()
desktopSessionRecording := desktopRecordingEnabled && userRoles.RecordDesktopSession()
directorySharing := userRoles.DesktopDirectorySharing()
download := newAccess(userRoles, ctx, types.KindDownload)
license := newAccess(userRoles, ctx, types.KindLicense)
deviceTrust := newAccess(userRoles, ctx, types.KindDevice)
integrationsAccess := newAccess(userRoles, ctx, types.KindIntegration)
discoveryConfigsAccess := newAccess(userRoles, ctx, types.KindDiscoveryConfig)
lockAccess := newAccess(userRoles, ctx, types.KindLock)
accessListAccess := newAccess(userRoles, ctx, types.KindAccessList)
externalAuditStorage := newAccess(userRoles, ctx, types.KindExternalAuditStorage)
bots := newAccess(userRoles, ctx, types.KindBot)
botInstances := newAccess(userRoles, ctx, types.KindBotInstance)
instances := newAccess(userRoles, ctx, types.KindInstance)
crownJewelAccess := newAccess(userRoles, ctx, types.KindCrownJewel)
userTasksAccess := newAccess(userRoles, ctx, types.KindUserTask)
reviewRequests := userRoles.MaybeCanReviewRequests()
fileTransferAccess := userRoles.CanCopyFiles()
workloadIdentity := newAccess(userRoles, ctx, types.KindWorkloadIdentity)
var auditQuery ResourceAccess
var securityReports ResourceAccess
if accessMonitoringEnabled {
auditQuery = newAccess(userRoles, ctx, types.KindAuditQuery)
securityReports = newAccess(userRoles, ctx, types.KindSecurityReport)
}
contact := newAccess(userRoles, ctx, types.KindContact)
var clientIPRestrictions ResourceAccess
if features.Cloud {
clientIPRestrictions = newAccess(userRoles, ctx, types.KindClientIPRestriction)
}
autoUpdateConfig := newAccess(userRoles, ctx, types.KindAutoUpdateConfig)
autoUpdateVersion := newAccess(userRoles, ctx, types.KindAutoUpdateVersion)
autoUpdateAgentRollout := newAccess(userRoles, ctx, types.KindAutoUpdateAgentRollout)
autoUpdateAgentReport := newAccess(userRoles, ctx, types.KindAutoUpdateAgentReport)
beamsEntitlement := modules.GetProtoEntitlement(&features, entitlements.Beams)
var beam ResourceAccess
if beamsEntitlement.Enabled {
beam = newAccess(userRoles, ctx, types.KindBeam)
}
mobileDevice := MobileDeviceAccess{
CreateEnrollToken: hasAccess(userRoles, ctx, types.KindMobileDevice, types.VerbCreateEnrollToken),
}
return UserACL{
AccessRequests: requestAccess,
AppServers: appServerAccess,
DBServers: dbServerAccess,
DB: dbAccess,
ReviewRequests: reviewRequests,
KubeServers: kubeServerAccess,
Desktops: desktopAccess,
AuthConnectors: authConnectors,
TrustedClusters: trustedClusterAccess,
RecordedSessions: recordedSessionAccess,
ActiveSessions: activeSessionAccess,
Roles: roleAccess,
Events: eventAccess,
Users: userAccess,
Tokens: tokenAccess,
Nodes: nodeAccess,
Billing: billingAccess,
ConnectionDiagnostic: cnDiagnosticAccess,
Clipboard: clipboard,
DesktopSessionRecording: desktopSessionRecording,
DirectorySharing: directorySharing,
Download: download,
License: license,
Plugins: pluginsAccess,
Integrations: integrationsAccess,
UserTasks: userTasksAccess,
DiscoveryConfig: discoveryConfigsAccess,
DeviceTrust: deviceTrust,
Locks: lockAccess,
SAMLIdpServiceProvider: samlIdpServiceProviderAccess,
AccessList: accessListAccess,
AuditQuery: auditQuery,
SecurityReport: securityReports,
ExternalAuditStorage: externalAuditStorage,
AccessGraph: accessGraphAccess,
Bots: bots,
BotInstances: botInstances,
Instances: instances,
AccessMonitoringRule: accessMonitoringRules,
CrownJewel: crownJewelAccess,
AccessGraphSettings: accessGraphSettings,
Contact: contact,
FileTransferAccess: fileTransferAccess,
WebTerminalClipboardMode: userRoles.GetWebTerminalClipboardMode(),
GitServers: gitServersAccess,
WorkloadIdentity: workloadIdentity,
ClientIPRestriction: clientIPRestrictions,
InferenceModel: newAccess(userRoles, ctx, types.KindInferenceModel),
InferencePolicy: newAccess(userRoles, ctx, types.KindInferencePolicy),
InferenceSecret: newAccess(userRoles, ctx, types.KindInferenceSecret),
Classifier: newAccess(userRoles, ctx, types.KindClassifier),
AutoUpdateConfig: autoUpdateConfig,
AutoUpdateVersion: autoUpdateVersion,
AutoUpdateAgentRollout: autoUpdateAgentRollout,
AutoUpdateAgentReport: autoUpdateAgentReport,
Beam: beam,
MobileDevice: mobileDevice,
}
}
/*
* 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/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UserGroups defines an interface for managing UserGroups.
type UserGroups interface {
// ListUserGroups returns a paginated list of all user group resources.
ListUserGroups(context.Context, int, string) ([]types.UserGroup, string, error)
// GetUserGroup returns the specified user group resources.
GetUserGroup(ctx context.Context, name string) (types.UserGroup, error)
// CreateUserGroup creates a new user group resource.
CreateUserGroup(context.Context, types.UserGroup) error
// UpdateUserGroup updates an existing user group resource.
UpdateUserGroup(context.Context, types.UserGroup) error
// DeleteUserGroup removes the specified user group resource.
DeleteUserGroup(ctx context.Context, name string) error
// DeleteAllUserGroups removes all user groups.
DeleteAllUserGroups(context.Context) error
}
// MarshalUserGroup marshals the user group resource to JSON.
func MarshalUserGroup(group types.UserGroup, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch g := group.(type) {
case *types.UserGroupV1:
if err := g.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return utils.FastMarshal(maybeResetProtoRevision(cfg.PreserveRevision, g))
default:
return nil, trace.BadParameter("unsupported user group resource %T", g)
}
}
// UnmarshalUserGroup unmarshals user group resource from JSON.
func UnmarshalUserGroup(data []byte, opts ...MarshalOption) (types.UserGroup, error) {
if len(data) == 0 {
return nil, trace.BadParameter("missing group data")
}
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
var h types.ResourceHeader
if err := utils.FastUnmarshal(data, &h); err != nil {
return nil, trace.Wrap(err)
}
switch h.Version {
case types.V1:
var g types.UserGroupV1
if err := utils.FastUnmarshal(data, &g); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := g.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if cfg.Revision != "" {
g.SetRevision(cfg.Revision)
}
if !cfg.Expires.IsZero() {
g.SetExpiry(cfg.Expires)
}
return &g, nil
}
return nil, trace.BadParameter("unsupported user group resource version %q", h.Version)
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalUserToken unmarshals the UserToken resource from JSON.
func UnmarshalUserToken(bytes []byte, opts ...MarshalOption) (types.UserToken, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
var token types.UserTokenV3
if err := utils.FastUnmarshal(bytes, &token); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := token.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &token, nil
}
// MarshalUserToken marshals the UserToken resource to JSON.
func MarshalUserToken(token types.UserToken, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch t := token.(type) {
case *types.UserTokenV3:
if err := t.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *t
copy.SetRevision("")
t = ©
}
return utils.FastMarshal(t)
default:
return nil, trace.BadParameter("unsupported user token resource %T", t)
}
}
/*
* 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/utils"
)
// UnmarshalUserTokenSecrets unmarshals the UserTokenSecrets resource from JSON.
func UnmarshalUserTokenSecrets(bytes []byte, opts ...MarshalOption) (types.UserTokenSecrets, error) {
if len(bytes) == 0 {
return nil, trace.BadParameter("missing resource data")
}
var secrets types.UserTokenSecretsV3
if err := utils.FastUnmarshal(bytes, &secrets); err != nil {
return nil, trace.BadParameter("%s", err)
}
if err := secrets.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &secrets, nil
}
// MarshalUserTokenSecrets marshals the UserTokenSecrets resource to JSON.
func MarshalUserTokenSecrets(secrets types.UserTokenSecrets, opts ...MarshalOption) ([]byte, error) {
cfg, err := CollectOptions(opts)
if err != nil {
return nil, trace.Wrap(err)
}
switch t := secrets.(type) {
case *types.UserTokenSecretsV3:
if err := t.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
if !cfg.PreserveRevision {
copy := *t
copy.SetRevision("")
t = ©
}
return utils.FastMarshal(t)
default:
return nil, trace.BadParameter("unsupported user token secrets resource %T", t)
}
}
/*
* 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"
"log/slog"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
healthcheckconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/healthcheckconfig/v1"
presencev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/presence/v1"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/api/utils/retryutils"
"github.com/gravitational/teleport/lib/defaults"
iterstream "github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/scopes"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
const (
// smallFanoutCapacity is the default capacity used for the circular event buffer allocated by
// resource watchers that implement event fanout.
smallFanoutCapacity = 128
// eventBufferMaxSize is the maximum size of the event buffer used by resource watchers to
// batch events that arrive in quick succession. In practice the event buffer should never
// grow this large unless we're dealing with a truly massive teleport cluster.
eventBufferMaxSize = 2048
)
// resourceCollector is a generic interface for maintaining an up-to-date view
// of a resource set being monitored. Used in conjunction with resourceWatcher.
type resourceCollector interface {
// resourceKinds specifies the resource kind to watch.
resourceKinds() []types.WatchKind
// getResourcesAndUpdateCurrent is called when the resources should be
// (re-)fetched directly.
getResourcesAndUpdateCurrent(context.Context) error
// processEventsAndUpdateCurrent is called when a watcher events are received. The event buffer
// may be reused so implementers must not retain it, but implementers may mutate the buffer
// in place during the call, e.g. in order to filter out undesired events before passing them
// to a subsideary bulk-processor such as a fanout.
processEventsAndUpdateCurrent(context.Context, []types.Event)
// notifyStale is called when the maximum acceptable staleness (if specified)
// is exceeded.
notifyStale()
// initializationChan is used to check if the initial state sync has
// been completed.
initializationChan() <-chan struct{}
}
func watchKindsString(kinds []types.WatchKind) string {
var sb strings.Builder
for i, k := range kinds {
if i != 0 {
sb.WriteString(", ")
}
sb.WriteString(k.Kind)
if k.SubKind != "" {
sb.WriteString("/")
sb.WriteString(k.SubKind)
}
}
return sb.String()
}
// ResourceWatcherConfig configures resource watcher.
type ResourceWatcherConfig struct {
// Clock is used to control time.
Clock clockwork.Clock
// Client is used to create new watchers
Client types.Events
// Logger emits log messages.
Logger *slog.Logger
// ResetC is a channel to notify of internal watcher reset (used in tests).
ResetC chan time.Duration
// Component is a component used in logs.
Component string
// MaxRetryPeriod is the maximum retry period on failed watchers.
MaxRetryPeriod time.Duration
// MaxStaleness is a maximum acceptable staleness for the locally maintained
// resources, zero implies no staleness detection.
MaxStaleness time.Duration
// QueueSize is an optional queue size
QueueSize int
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *ResourceWatcherConfig) CheckAndSetDefaults() error {
if cfg.Component == "" {
return trace.BadParameter("missing parameter Component")
}
if cfg.Logger == nil {
cfg.Logger = slog.Default()
}
if cfg.MaxRetryPeriod == 0 {
cfg.MaxRetryPeriod = defaults.MaxWatcherBackoff
}
if cfg.Clock == nil {
cfg.Clock = clockwork.NewRealClock()
}
if cfg.Client == nil {
return trace.BadParameter("missing parameter Client")
}
if cfg.ResetC == nil {
cfg.ResetC = make(chan time.Duration, 1)
}
return nil
}
// newResourceWatcher returns a new instance of resourceWatcher.
// It is the caller's responsibility to verify the inputs' validity
// incl. cfg.CheckAndSetDefaults.
func newResourceWatcher(ctx context.Context, collector resourceCollector, cfg ResourceWatcherConfig) (*resourceWatcher, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
retry, err := retryutils.NewLinear(retryutils.LinearConfig{
First: retryutils.FullJitter(cfg.MaxRetryPeriod / 10),
Step: cfg.MaxRetryPeriod / 5,
Max: cfg.MaxRetryPeriod,
Jitter: retryutils.HalfJitter,
Clock: cfg.Clock,
})
if err != nil {
return nil, trace.Wrap(err)
}
cfg.Logger = cfg.Logger.With("resource_kinds", watchKindsString(collector.resourceKinds()))
ctx, cancel := context.WithCancel(ctx)
p := &resourceWatcher{
ResourceWatcherConfig: cfg,
collector: collector,
ctx: ctx,
cancel: cancel,
retry: retry,
LoopC: make(chan struct{}),
StaleC: make(chan struct{}),
}
go p.runWatchLoop()
return p, nil
}
// resourceWatcher monitors additions, updates and deletions
// to a set of resources.
type resourceWatcher struct {
// failureStartedAt records when the current sync failures were first
// detected, zero if there are no failures present.
failureStartedAt time.Time
collector resourceCollector
// ctx is a context controlling the lifetime of this resourceWatcher
// instance.
ctx context.Context
// retry is used to manage backoff logic for watchers.
retry retryutils.Retry
cancel context.CancelFunc
// LoopC is a channel to check whether the watch loop is running
// (used in tests).
LoopC chan struct{}
// StaleC is a channel that can trigger the condition of resource staleness
// (used in tests).
StaleC chan struct{}
ResourceWatcherConfig
}
// Done returns a channel that signals resource watcher closure.
func (p *resourceWatcher) Done() <-chan struct{} {
return p.ctx.Done()
}
// Close closes the resource watcher and cancels all the functions.
func (p *resourceWatcher) Close() {
p.cancel()
}
// IsInitialized is a non-blocking way to check if resource watcher is already
// initialized.
func (p *resourceWatcher) IsInitialized() bool {
select {
case <-p.collector.initializationChan():
return true
default:
return false
}
}
// WaitInitialization blocks until resource watcher is fully initialized with
// the resources presented in auth server.
func (p *resourceWatcher) WaitInitialization() error {
// wait for resourceWatcher to complete initialization.
t := time.NewTicker(5 * time.Second)
defer t.Stop()
for {
select {
case <-p.collector.initializationChan():
return nil
case <-t.C:
p.Logger.DebugContext(p.ctx, "ResourceWatcher is not yet initialized.")
case <-p.ctx.Done():
return trace.BadParameter("ResourceWatcher %s failed to initialize.", watchKindsString(p.collector.resourceKinds()))
}
}
}
// hasStaleView returns true when the local view has failed to be updated
// for longer than the MaxStaleness bound.
func (p *resourceWatcher) hasStaleView() bool {
// Used for testing stale lock views.
select {
case <-p.StaleC:
return true
default:
}
if p.MaxStaleness == 0 || p.failureStartedAt.IsZero() {
return false
}
return p.Clock.Since(p.failureStartedAt) > p.MaxStaleness
}
// runWatchLoop runs a watch loop.
func (p *resourceWatcher) runWatchLoop() {
for {
p.Logger.Log(p.ctx, logutils.TraceLevel, "Starting watch.")
err := p.watch()
select {
case <-p.ctx.Done():
return
default:
}
if err != nil && p.failureStartedAt.IsZero() {
// Note that failureStartedAt is zeroed in the watch routine immediately
// after the local resource set has been successfully updated.
p.failureStartedAt = p.Clock.Now()
}
if p.hasStaleView() {
p.Logger.WarnContext(p.ctx, "Maximum staleness of period exceeded.", "max_staleness", p.MaxStaleness, "failure_started", p.failureStartedAt)
p.collector.notifyStale()
}
// Used for testing that the watch routine has exited and is about
// to be restarted.
select {
case p.ResetC <- p.retry.Duration():
default:
}
startedWaiting := p.Clock.Now()
select {
case t := <-p.retry.After():
p.Logger.DebugContext(p.ctx, "Attempting to restart watch after waiting", "waited", t.Sub(startedWaiting))
p.retry.Inc()
case <-p.ctx.Done():
p.Logger.DebugContext(p.ctx, "Closed, returning from watch loop.")
return
case <-p.StaleC:
// Used for testing that the watch routine is waiting for the
// next restart attempt. We don't want to wait for the full
// retry period in tests so we trigger the restart immediately.
p.Logger.DebugContext(p.ctx, "Stale view, continue watch loop.")
}
if err != nil {
p.Logger.WarnContext(p.ctx, "Restart watch on error", "error", err)
}
}
}
// watch monitors new resource updates, maintains a local view and broadcasts
// notifications to connected agents.
func (p *resourceWatcher) watch() error {
watch := types.Watch{
Name: p.Component,
MetricComponent: p.Component,
Kinds: p.collector.resourceKinds(),
}
if p.QueueSize > 0 {
watch.QueueSize = p.QueueSize
}
watcher, err := p.Client.NewWatcher(p.ctx, watch)
if err != nil {
return trace.Wrap(err)
}
defer watcher.Close()
// before fetch, make sure watcher is synced by receiving init event,
// to avoid the scenario:
// 1. Cache process: w = NewWatcher()
// 2. Cache process: c.fetch()
// 3. Backend process: addItem()
// 4. Cache process: <- w.Events()
//
// If there is a way that NewWatcher() on line 1 could
// return without subscription established first,
// Code line 3 could execute and line 4 could miss event,
// wrapping up with out of sync replica.
// To avoid this, before doing fetch,
// cache process makes sure the connection is established
// by receiving init event first.
select {
case <-watcher.Done():
return trace.ConnectionProblem(watcher.Error(), "watcher is closed: %v", watcher.Error())
case <-p.ctx.Done():
return trace.ConnectionProblem(p.ctx.Err(), "context is closing")
case <-p.StaleC:
return trace.ConnectionProblem(nil, "stale view")
case event := <-watcher.Events():
if event.Type != types.OpInit {
return trace.BadParameter("expected init event, got %v instead", event.Type)
}
}
if err := p.collector.getResourcesAndUpdateCurrent(p.ctx); err != nil {
return trace.Wrap(err)
}
p.retry.Reset()
p.failureStartedAt = time.Time{}
// start out with a modestly sized event buffer
eventBuf := make([]types.Event, 0, 16)
for {
select {
case <-watcher.Done():
return trace.ConnectionProblem(watcher.Error(), "watcher is closed: %v", watcher.Error())
case <-p.ctx.Done():
return trace.ConnectionProblem(p.ctx.Err(), "context is closing")
case event := <-watcher.Events():
// resource collectors want to process events in batches
// when possible in order to reduce contention on their locks.
// we therefore optimistically try to gather a large number of
// events without blocking.
eventBuf = append(eventBuf, event)
CollectEvents:
for len(eventBuf) < eventBufferMaxSize {
select {
case additionalEvent := <-watcher.Events():
eventBuf = append(eventBuf, additionalEvent)
default:
break CollectEvents
}
}
p.collector.processEventsAndUpdateCurrent(p.ctx, eventBuf)
clear(eventBuf)
eventBuf = eventBuf[:0]
case p.LoopC <- struct{}{}:
// Used in tests to detect the watch loop is running.
case <-p.StaleC:
return trace.ConnectionProblem(nil, "stale view")
}
}
}
// ProxyWatcherConfig is a ProxyWatcher configuration.
type ProxyWatcherConfig struct {
// ProxyGetter is used to directly fetch the list of active proxies.
ProxyGetter
// ProxyDiffer is used to decide whether a put operation on an existing proxy should
// trigger a event.
ProxyDiffer func(old, new types.Server) bool
// ProxiesC is a channel used to report the current proxy set. It receives
// a fresh list at startup and subsequently a list of all known proxy
// whenever an addition or deletion is detected.
ProxiesC chan []types.Server
ResourceWatcherConfig
}
// NewProxyWatcher returns a new instance of GenericWatcher that is configured
// to watch for changes.
func NewProxyWatcher(ctx context.Context, cfg ProxyWatcherConfig) (*GenericWatcher[types.Server, readonly.Server], error) {
if cfg.ProxyGetter == nil {
return nil, trace.BadParameter("ProxyGetter must be provided")
}
if cfg.ProxyDiffer == nil {
cfg.ProxyDiffer = func(old, new types.Server) bool { return true }
}
proxyGetter := cfg.ProxyGetter
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.Server, readonly.Server]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindProxy,
ResourceKey: types.Server.GetName,
ResourceGetter: func(ctx context.Context) ([]types.Server, error) {
return clientutils.CollectWithFallback(ctx, proxyGetter.ListProxyServers, func(context.Context) ([]types.Server, error) {
//nolint:staticcheck // TODO(kiosion) DELETE IN 21.0.0
return proxyGetter.GetProxies()
})
},
ResourcesC: cfg.ProxiesC,
ResourceDiffer: cfg.ProxyDiffer,
RequireResourcesForInitialBroadcast: true,
CloneFunc: types.Server.DeepCopy,
ReadOnlyFunc: func(resource types.Server) readonly.Server {
return resource
},
})
return w, trace.Wrap(err)
}
// DatabaseWatcherConfig is a DatabaseWatcher configuration.
type DatabaseWatcherConfig struct {
// DatabaseGetter is responsible for fetching database resources.
DatabaseGetter
// DatabasesC receives up-to-date list of all database resources.
DatabasesC chan []types.Database
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
}
// NewDatabaseWatcher returns a new instance of DatabaseWatcher.
func NewDatabaseWatcher(ctx context.Context, cfg DatabaseWatcherConfig) (*GenericWatcher[types.Database, readonly.Database], error) {
if cfg.DatabaseGetter == nil {
return nil, trace.BadParameter("DatabaseGetter must be provided")
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.Database, readonly.Database]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindDatabase,
ResourceKey: types.Database.GetName,
ResourceGetter: cfg.DatabaseGetter.GetDatabases,
ResourcesC: cfg.DatabasesC,
CloneFunc: func(resource types.Database) types.Database {
return resource.Copy()
},
ReadOnlyFunc: func(resource types.Database) readonly.Database {
return resource
},
})
return w, trace.Wrap(err)
}
// AppWatcherConfig is an AppWatcher configuration.
type AppWatcherConfig struct {
// AppGetter is responsible for fetching application resources.
AppGetter
// AppsC receives up-to-date list of all application resources.
AppsC chan []types.Application
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
}
// NewAppWatcher returns a new instance of AppWatcher.
func NewAppWatcher(ctx context.Context, cfg AppWatcherConfig) (*GenericWatcher[types.Application, readonly.Application], error) {
if cfg.AppGetter == nil {
return nil, trace.BadParameter("AppGetter must be provided")
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.Application, readonly.Application]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindApp,
ResourceKey: types.Application.GetName,
ResourceGetter: cfg.AppGetter.GetApps,
ResourcesC: cfg.AppsC,
CloneFunc: func(resource types.Application) types.Application {
return resource.Copy()
},
ReadOnlyFunc: func(resource types.Application) readonly.Application {
return resource
},
})
return w, trace.Wrap(err)
}
type AppServersWatcherConfig struct {
AppServersGetter
ResourceWatcherConfig
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *AppServersWatcherConfig) CheckAndSetDefaults() error {
if err := cfg.ResourceWatcherConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if cfg.MaxStaleness == 0 {
const appServerMaxStaleness = time.Minute
cfg.MaxStaleness = appServerMaxStaleness
}
if cfg.AppServersGetter == nil {
getter, ok := cfg.Client.(AppServersGetter)
if !ok {
return trace.BadParameter("missing parameter AppServersGetter and Client not usable as AppServersGetter")
}
cfg.AppServersGetter = getter
}
return nil
}
func NewAppServersWatcher(ctx context.Context, cfg AppServersWatcherConfig) (*GenericWatcher[types.AppServer, readonly.AppServer], error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.AppServer, readonly.AppServer]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindAppServer,
// The app server watcher is a proxy only routing construct (see the
// reversetunnel server and initProxyEndpoint call sites); it must observe app
// servers in every scope, not just unscoped ones, so it watches with MODE_ALL.
ScopeFilter: types.ScopeFilterFromProto(scopesv1.Filter_builder{Mode: scopesv1.Mode_MODE_ALL}.Build()),
ResourceKey: func(resource types.AppServer) string {
// host IDs are guaranteed to not contain "/"
return resource.GetHostID() + "/" + resource.GetName()
},
DeleteKey: func(r types.Resource) string {
// the host ID is stored in metadata.description in app server delete events
return r.GetMetadata().Description + "/" + r.GetName()
},
ResourceGetter: func(ctx context.Context) ([]types.AppServer, error) {
// TODO(fspmarshall/scopes): this list does not yet honor scope filters, so it
// currently returns app servers in every scope and happens to match the MODE_ALL watch
// above. Once the list API supports scope filters (and defaults unscoped callers to
// unscoped-only, like the watch API), this call must explicitly request MODE_ALL as well.
return cfg.AppServersGetter.GetApplicationServers(ctx, apidefaults.Namespace)
},
DisableUpdateBroadcast: true,
CloneFunc: types.AppServer.Copy,
ReadOnlyFunc: func(resource types.AppServer) readonly.AppServer {
return resource
},
})
return w, trace.Wrap(err)
}
type DatabaseServerWatcherConfig struct {
DatabaseServersGetter
ResourceWatcherConfig
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *DatabaseServerWatcherConfig) CheckAndSetDefaults() error {
if err := cfg.ResourceWatcherConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if cfg.MaxStaleness == 0 {
const databaseServerMaxStaleness = time.Minute
cfg.MaxStaleness = databaseServerMaxStaleness
}
if cfg.DatabaseServersGetter == nil {
getter, ok := cfg.Client.(DatabaseServersGetter)
if !ok {
return trace.BadParameter("missing parameter DatabaseServersGetter and Client not usable as DatabaseServersGetter")
}
cfg.DatabaseServersGetter = getter
}
return nil
}
func NewDatabaseServerWatcher(ctx context.Context, cfg DatabaseServerWatcherConfig) (*GenericWatcher[types.DatabaseServer, readonly.DatabaseServer], error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.DatabaseServer, readonly.DatabaseServer]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindDatabaseServer,
ResourceKey: func(r types.DatabaseServer) string {
// the host ID is guaranteed not to contain "/"
return r.GetHostID() + "/" + r.GetName()
},
DeleteKey: func(r types.Resource) string {
// database servers put the host ID in the description in delete events
return r.GetMetadata().Description + "/" + r.GetName()
},
ResourceGetter: func(ctx context.Context) ([]types.DatabaseServer, error) {
return cfg.DatabaseServersGetter.GetDatabaseServers(ctx, apidefaults.Namespace)
},
DisableUpdateBroadcast: true,
CloneFunc: types.DatabaseServer.Copy,
ReadOnlyFunc: readonly.NewDatabaseServer,
})
return w, trace.Wrap(err)
}
// KubeServerWatcherConfig is an KubeServerWatcher configuration.
type KubeServerWatcherConfig struct {
// KubernetesServerGetter is responsible for fetching kube_server resources.
KubernetesServerGetter
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
}
// NewKubeServerWatcher returns a new instance of KubeServerWatcher.
func NewKubeServerWatcher(ctx context.Context, cfg KubeServerWatcherConfig) (*GenericWatcher[types.KubeServer, readonly.KubeServer], error) {
if cfg.KubernetesServerGetter == nil {
return nil, trace.BadParameter("KubernetesServerGetter must be provided")
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.KubeServer, readonly.KubeServer]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindKubeServer,
ResourceGetter: cfg.KubernetesServerGetter.GetKubernetesServers,
ResourceKey: func(resource types.KubeServer) string {
return resource.GetHostID() + resource.GetName()
},
DisableUpdateBroadcast: true,
CloneFunc: types.KubeServer.Copy,
ReadOnlyFunc: func(resource types.KubeServer) readonly.KubeServer {
return resource
},
})
return w, trace.Wrap(err)
}
// KubeClusterWatcherConfig is an KubeClusterWatcher configuration.
type KubeClusterWatcherConfig struct {
// KubernetesGetter is responsible for fetching kube_cluster resources.
KubernetesClusterGetter
// KubeClustersC receives up-to-date list of all kube_cluster resources.
KubeClustersC chan []types.KubeCluster
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
// LoadSecrets specifies whether the watched kube clusters include their kubeconfig. Only the
// kube agent needs this, to connect to dynamically-registered clusters, and it requires
// secret-inclusive read permission on kubernetes_cluster.
LoadSecrets bool
}
// NewKubeClusterWatcher returns a new instance of KubeClusterWatcher.
func NewKubeClusterWatcher(ctx context.Context, cfg KubeClusterWatcherConfig) (*GenericWatcher[types.KubeCluster, readonly.KubeCluster], error) {
if cfg.KubernetesClusterGetter == nil {
return nil, trace.BadParameter("KubernetesClusterGetter must be provided")
}
getter := cfg.KubernetesClusterGetter
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.KubeCluster, readonly.KubeCluster]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindKubernetesCluster,
LoadSecrets: cfg.LoadSecrets,
ResourceGetter: func(ctx context.Context) ([]types.KubeCluster, error) {
return iterstream.Collect(getter.RangeKubeClusters(ctx, presencev1.ListKubeClustersRequest_builder{
WithSecrets: cfg.LoadSecrets,
}.Build()))
},
ResourceKey: GetCursorForKubeCluster,
DeleteKey: func(res types.Resource) string {
cluster, ok := res.(types.KubeCluster)
if !ok {
return ""
}
return GetCursorForKubeCluster(cluster)
},
ResourcesC: cfg.KubeClustersC,
CloneFunc: func(resource types.KubeCluster) types.KubeCluster {
return resource.Copy()
},
ReadOnlyFunc: func(resource types.KubeCluster) readonly.KubeCluster {
return resource
},
})
return w, trace.Wrap(err)
}
type DynamicWindowsDesktopGetter interface {
ListDynamicWindowsDesktops(ctx context.Context, pageSize int, pageToken string) ([]types.DynamicWindowsDesktop, string, error)
}
// DynamicWindowsDesktopWatcherConfig is a DynamicWindowsDesktopWatcher configuration.
type DynamicWindowsDesktopWatcherConfig struct {
// DynamicWindowsDesktopGetter is responsible for fetching DynamicWindowsDesktop resources.
DynamicWindowsDesktopGetter
// DynamicWindowsDesktopsC receives up-to-date list of all DynamicWindowsDesktop resources.
DynamicWindowsDesktopsC chan []types.DynamicWindowsDesktop
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
}
// NewDynamicWindowsDesktopWatcher returns a new instance of DynamicWindowsDesktopWatcher.
func NewDynamicWindowsDesktopWatcher(ctx context.Context, cfg DynamicWindowsDesktopWatcherConfig) (*GenericWatcher[types.DynamicWindowsDesktop, readonly.DynamicWindowsDesktop], error) {
if cfg.DynamicWindowsDesktopGetter == nil {
return nil, trace.BadParameter("DynamicWindowsDesktopGetter must be provided")
}
getter := cfg.DynamicWindowsDesktopGetter
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.DynamicWindowsDesktop, readonly.DynamicWindowsDesktop]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindDynamicWindowsDesktop,
ResourceGetter: pagerFn[types.DynamicWindowsDesktop](
getter.ListDynamicWindowsDesktops,
).getAll,
ResourceKey: types.DynamicWindowsDesktop.GetName,
ResourcesC: cfg.DynamicWindowsDesktopsC,
CloneFunc: types.DynamicWindowsDesktop.Copy,
ReadOnlyFunc: func(resource types.DynamicWindowsDesktop) readonly.DynamicWindowsDesktop {
return resource
},
})
return w, trace.Wrap(err)
}
// GenericWatcherConfig is a generic resource watcher configuration.
type GenericWatcherConfig[T any, R any] struct {
// ResourceGetter is used to directly fetch the current set of resources.
ResourceGetter func(context.Context) ([]T, error)
// ResourceDiffer is used to decide whether a put operation on an existing ResourceGetter should
// trigger an event.
ResourceDiffer func(old, new T) bool
// ResourceKey defines how the resources should be keyed.
ResourceKey func(resource T) string
// DeleteKey defines how a deleted resource key is derived. A delete event
// typically sends a stripped down resource representation with an underlying
// type of [types.ResourceHeader].
// If unspecified the key will be derived from the resource.Description + resource.GetName
DeleteKey func(types.Resource) string
// ResourcesC is a channel used to report the current resource set. It receives
// a fresh list at startup and subsequently a list of all known resources
// whenever an addition or deletion is detected.
ResourcesC chan []T
// CloneFunc defines how a resource is cloned. All resources provided via
// the broadcast mechanism, or retrieved via [GenericWatcer.CurrentResources]
// or [GenericWatcher.CurrentResourcesWithFilter] will be cloned by this
// mechanism before being provided to callers.
CloneFunc func(resource T) T
// ReadOnlyFunc returns the read-only view of a resource. Ideally this will
// be a type conversion (but we can't statically express that as constraints
// on T and R) but it's also possible to wrapper the original resource.
// Making the read-only view should be much cheaper than CloneFunc.
ReadOnlyFunc func(resource T) R
ResourceWatcherConfig
// ResourceKind specifies the kind of resource the watcher is monitoring.
ResourceKind string
// ResourceFilter is an optional filter that is applied on the backend when
// watching for resources. Only resources matching the filter will be sent
// to the watcher.
ResourceFilter map[string]string
// ScopeFilter is an optional scope filter applied to the watch. A nil filter
// yields the caller's default scope behavior (unscoped-only for unscoped
// callers, current-scope-only for scoped callers). Watchers that run on the
// teleport proxy and must observe every instance of an optionally-scoped kind
// (regardless of scope) should set this to MODE_ALL. Note that the paired
// ResourceGetter (the list seed) does not yet honor scope filters; see the
// TODO at each MODE_ALL call site.
ScopeFilter *types.ScopeFilter
// RequireResourcesForInitialBroadcast indicates whether an update should be
// performed if the initial set of resources is empty.
RequireResourcesForInitialBroadcast bool
// DisableUpdateBroadcast turns off emitting updates on changes. When this
// mode is opted into, users must invoke [GenericWatcher.CurrentResources] or
// [GenericWatcher.CurrentResourcesWithFilter] manually to retrieve the active
// resource set.
DisableUpdateBroadcast bool
// LoadSecrets specifies whether sensitive data will be loaded into memory.
// This is only applicable to certain types like [types.CertAuthority].
LoadSecrets bool
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *GenericWatcherConfig[T, R]) CheckAndSetDefaults() error {
if err := cfg.ResourceWatcherConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if cfg.ResourceGetter == nil {
return trace.BadParameter("ResourceGetter not provided to generic resource watcher")
}
if cfg.ResourceKind == "" {
return trace.BadParameter("ResourceKind not provided to generic resource watcher")
}
if cfg.ResourceKey == nil {
return trace.BadParameter("ResourceKey not provided to generic resource watcher")
}
if cfg.CloneFunc == nil {
return trace.BadParameter("CloneFunc not provided to generic resource watcher")
}
if cfg.ReadOnlyFunc == nil {
return trace.BadParameter("ReadOnlyFunc not provided to generic resource watcher")
}
if cfg.ResourceDiffer == nil {
cfg.ResourceDiffer = func(T, T) bool { return true }
}
if cfg.ResourcesC == nil {
cfg.ResourcesC = make(chan []T)
}
return nil
}
// NewGenericResourceWatcher returns a new instance of resource watcher.
func NewGenericResourceWatcher[T any, R any](ctx context.Context, cfg GenericWatcherConfig[T, R]) (*GenericWatcher[T, R], error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cache, err := utils.NewFnCache(utils.FnCacheConfig{
Context: ctx,
TTL: 3 * time.Second,
Clock: cfg.Clock,
})
if err != nil {
return nil, trace.Wrap(err)
}
collector := &genericCollector[T, R]{
GenericWatcherConfig: cfg,
initializationC: make(chan struct{}),
cache: cache,
}
collector.stale.Store(true)
watcher, err := newResourceWatcher(ctx, collector, cfg.ResourceWatcherConfig)
if err != nil {
return nil, trace.Wrap(err)
}
return &GenericWatcher[T, R]{watcher, collector}, nil
}
// GenericWatcher is built on top of resourceWatcher to monitor additions
// and deletions to the set of resources.
type GenericWatcher[T any, R any] struct {
*resourceWatcher
*genericCollector[T, R]
}
// ResourceCount returns the current number of resources known to the watcher.
func (g *GenericWatcher[T, R]) ResourceCount() int {
g.rw.RLock()
defer g.rw.RUnlock()
return len(g.current)
}
// CurrentResources returns a copy of the resources known to the watcher.
func (g *GenericWatcher[T, R]) CurrentResources(ctx context.Context) ([]T, error) {
if err := g.refreshStaleResources(ctx); err != nil {
return nil, trace.Wrap(err)
}
g.rw.RLock()
defer g.rw.RUnlock()
return resourcesToSlice(g.current, g.CloneFunc), nil
}
// CurrentResourcesWithFilter returns a copy of the resources known to the watcher
// that match the provided filter.
func (g *GenericWatcher[T, R]) CurrentResourcesWithFilter(ctx context.Context, filter func(R) bool) ([]T, error) {
if err := g.refreshStaleResources(ctx); err != nil {
return nil, trace.Wrap(err)
}
g.rw.RLock()
defer g.rw.RUnlock()
var out []T
for _, resource := range g.current {
if filter(g.ReadOnlyFunc(resource)) {
out = append(out, g.CloneFunc(resource))
}
}
return out, nil
}
// genericCollector accompanies resourceWatcher when monitoring proxies. T is
// the resource type, R is a read-only view over T.
type genericCollector[T any, R any] struct {
GenericWatcherConfig[T, R]
// current holds a map of the currently known resources (keyed by server name,
// RWMutex protected).
current map[string]T
initializationC chan struct{}
// cache is a helper for temporarily storing the results of CurrentResources.
// It's used to limit the number of calls to the backend.
cache *utils.FnCache
rw sync.RWMutex
once sync.Once
// stale is used to indicate that the watcher is stale and needs to be
// refreshed.
stale atomic.Bool
}
// resourceKinds specifies the resource kind to watch.
func (g *genericCollector[T, R]) resourceKinds() []types.WatchKind {
return []types.WatchKind{{
Kind: g.ResourceKind,
LoadSecrets: g.LoadSecrets,
Filter: g.ResourceFilter,
ScopeFilter: g.ScopeFilter,
}}
}
// getResources gets the list of current resources.
func (g *genericCollector[T, R]) getResources(ctx context.Context) (map[string]T, error) {
resources, err := g.GenericWatcherConfig.ResourceGetter(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
current := make(map[string]T, len(resources))
for _, resource := range resources {
current[g.GenericWatcherConfig.ResourceKey(resource)] = resource
}
return current, nil
}
func (g *genericCollector[T, R]) refreshStaleResources(ctx context.Context) error {
if !g.stale.Load() {
return nil
}
_, err := utils.FnCacheGet(ctx, g.cache, g.GenericWatcherConfig.ResourceKind, func(ctx context.Context) (any, error) {
newCurrent, err := g.getResources(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// as an optimization, we can check if the collector is still stale
// before grabbing the write lock
if !g.stale.Load() {
// the view is no longer stale, discard newCurrent and proceed with
// the data in g.current
return nil, nil
}
g.rw.Lock()
defer g.rw.Unlock()
// check the staleness again since it might've changed since we were not
// holding the lock
if !g.stale.Load() {
return nil, nil
}
g.current = newCurrent
g.stale.Store(false)
return nil, nil
})
return trace.Wrap(err)
}
// getResourcesAndUpdateCurrent is called when the resources should be
// (re-)fetched directly.
func (g *genericCollector[T, R]) getResourcesAndUpdateCurrent(ctx context.Context) error {
newCurrent, err := g.getResources(ctx)
if err != nil {
return trace.Wrap(err)
}
g.rw.Lock()
defer g.rw.Unlock()
g.current = newCurrent
g.stale.Store(false)
// Only emit an empty set of resources if the watcher is already initialized,
// or if explicitly opted into by for the watcher.
if len(newCurrent) > 0 || g.isInitialized() ||
(!g.RequireResourcesForInitialBroadcast && len(newCurrent) == 0) {
g.broadcastUpdate(ctx)
}
g.defineCollectorAsInitialized()
return nil
}
func (g *genericCollector[T, R]) defineCollectorAsInitialized() {
g.once.Do(func() {
// mark watcher as initialized.
close(g.initializationC)
})
}
// processEventsAndUpdateCurrent is called when a watcher event is received.
func (g *genericCollector[T, R]) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
g.rw.Lock()
defer g.rw.Unlock()
var updated bool
for _, event := range events {
if event.Resource == nil || event.Resource.GetKind() != g.ResourceKind {
g.Logger.WarnContext(ctx, "Received unexpected event", "event", logutils.StringerAttr(event))
continue
}
switch event.Type {
case types.OpDelete:
// On delete events, the server description is populated with the host ID.
key := event.Resource.GetMetadata().Description + event.Resource.GetName()
if g.DeleteKey != nil {
key = g.DeleteKey(event.Resource)
}
delete(g.current, key)
// Always broadcast when a resource is deleted.
updated = true
case types.OpPut:
resource, err := types.ConvertResource[T](event.Resource)
if err != nil {
g.Logger.WarnContext(ctx, "Failed to convert event resource",
"resource", event.Resource.GetKind(),
"error", err,
)
continue
}
key := g.ResourceKey(resource)
current, exists := g.current[key]
g.current[key] = resource
updated = !exists || g.ResourceDiffer(current, resource)
default:
g.Logger.WarnContext(ctx, "Skipping unsupported event type", "event_type", event.Type)
}
}
if updated {
g.broadcastUpdate(ctx)
}
}
// broadcastUpdate broadcasts information about updating the resource set.
func (g *genericCollector[T, R]) broadcastUpdate(ctx context.Context) {
if g.DisableUpdateBroadcast {
return
}
names := make([]string, 0, len(g.current))
for k := range g.current {
names = append(names, k)
}
g.Logger.DebugContext(ctx, "List of known resources updated", "resources", names)
select {
case g.ResourcesC <- resourcesToSlice(g.current, g.CloneFunc):
case <-ctx.Done():
}
}
// isInitialized is used to check that the cache has done its initial
// sync
func (g *genericCollector[T, R]) initializationChan() <-chan struct{} {
return g.initializationC
}
func (g *genericCollector[T, R]) isInitialized() bool {
select {
case <-g.initializationC:
return true
default:
return false
}
}
func (g *genericCollector[T, R]) notifyStale() {
g.stale.Store(true)
}
// LockWatcherConfig is a LockWatcher configuration.
type LockWatcherConfig struct {
LockGetter
ResourceWatcherConfig
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *LockWatcherConfig) CheckAndSetDefaults() error {
if err := cfg.ResourceWatcherConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if cfg.MaxStaleness == 0 {
cfg.MaxStaleness = defaults.LockMaxStaleness
}
if cfg.LockGetter == nil {
getter, ok := cfg.Client.(LockGetter)
if !ok {
return trace.BadParameter("missing parameter LockGetter and Client not usable as LockGetter")
}
cfg.LockGetter = getter
}
return nil
}
// NewLockWatcher returns a new instance of LockWatcher.
func NewLockWatcher(ctx context.Context, cfg LockWatcherConfig) (*LockWatcher, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
collector := &lockCollector{
LockWatcherConfig: cfg,
fanout: NewFanoutV2(FanoutV2Config{
Capacity: smallFanoutCapacity,
}),
initializationC: make(chan struct{}),
}
// Resource watcher require the fanout to be initialized before passing in.
// Otherwise, Emit() may fail due to a race condition mentioned in https://github.com/gravitational/teleport/issues/19289
collector.fanout.SetInit(collector.resourceKinds())
watcher, err := newResourceWatcher(ctx, collector, cfg.ResourceWatcherConfig)
if err != nil {
return nil, trace.Wrap(err)
}
return &LockWatcher{watcher, collector}, nil
}
// LockWatcher is built on top of resourceWatcher to monitor changes to locks.
type LockWatcher struct {
*resourceWatcher
*lockCollector
}
// lockCollector accompanies resourceWatcher when monitoring locks.
type lockCollector struct {
LockWatcherConfig
// current holds a map of the currently known locks (keyed by lock name).
current map[string]types.Lock
// fanout provides support for multiple subscribers to the lock updates.
fanout *FanoutV2
// initializationC is used to check whether the initial sync has completed
initializationC chan struct{}
// currentRW is a mutex protecting both current and isStale.
currentRW sync.RWMutex
once sync.Once
// isStale indicates whether the local lock view (current) is stale.
isStale bool
}
// IsStale is used to check whether the lock watcher is stale.
// Used in tests.
func (p *lockCollector) IsStale() bool {
p.currentRW.RLock()
defer p.currentRW.RUnlock()
return p.isStale
}
// Subscribe is used to subscribe to the lock updates.
func (p *lockCollector) Subscribe(ctx context.Context, targets ...types.LockTarget) (types.Watcher, error) {
watchKinds, err := lockTargetsToWatchKinds(targets)
if err != nil {
return nil, trace.Wrap(err)
}
sub, err := p.fanout.NewWatcher(ctx, types.Watch{Kinds: watchKinds})
if err != nil {
return nil, trace.Wrap(err)
}
select {
case event := <-sub.Events():
if event.Type != types.OpInit {
return nil, trace.BadParameter("expected init event, got %v instead", event.Type)
}
case <-sub.Done():
return nil, trace.Wrap(sub.Error())
}
return sub, nil
}
// CheckLockInForce returns an AccessDenied error if there is a lock in force
// matching at least one of the targets.
func (p *lockCollector) CheckLockInForce(mode constants.LockingMode, targets ...types.LockTarget) error {
if len(targets) == 0 {
// A lock can't match any targets if there are no targets.
return nil
}
p.currentRW.RLock()
defer p.currentRW.RUnlock()
if p.isStale && mode == constants.LockingModeStrict {
return StrictLockingModeAccessDenied
}
if lock := p.findLockInForceUnderMutex(targets); lock != nil {
return LockInForceAccessDenied(lock)
}
return nil
}
func (p *lockCollector) findLockInForceUnderMutex(targets []types.LockTarget) types.Lock {
for _, lock := range p.current {
if !lock.IsInForce(p.Clock.Now()) {
continue
}
for _, target := range targets {
if target.Match(lock) {
return lock
}
}
}
return nil
}
// GetCurrent returns the currently stored locks.
func (p *lockCollector) GetCurrent() []types.Lock {
p.currentRW.RLock()
defer p.currentRW.RUnlock()
return lockMapValues(p.current)
}
// resourceKinds specifies the resource kind to watch.
func (p *lockCollector) resourceKinds() []types.WatchKind {
return []types.WatchKind{{Kind: types.KindLock}}
}
// initializationChan is used to check that the cache has done its initial
// sync
func (p *lockCollector) initializationChan() <-chan struct{} {
return p.initializationC
}
// getResourcesAndUpdateCurrent is called when the resources should be
// (re-)fetched directly.
func (p *lockCollector) getResourcesAndUpdateCurrent(ctx context.Context) error {
locks, err := clientutils.CollectWithFallback(
ctx,
func(ctx context.Context, limit int, start string) ([]types.Lock, string, error) {
return p.LockGetter.ListLocks(ctx, limit, start, &types.LockFilter{InForceOnly: true})
},
func(ctx context.Context) ([]types.Lock, error) {
// TODO(okraport): DELETE IN v21
const inForceOnlyTrue = true
return p.LockGetter.GetLocks(ctx, inForceOnlyTrue)
},
)
if err != nil {
return trace.Wrap(err)
}
newCurrent := map[string]types.Lock{}
for _, lock := range locks {
newCurrent[lock.GetName()] = lock
}
p.currentRW.Lock()
defer p.currentRW.Unlock()
p.current = newCurrent
p.isStale = false
p.defineCollectorAsInitialized()
for _, lock := range p.current {
p.fanout.Emit(types.Event{Type: types.OpPut, Resource: lock})
}
return nil
}
func (p *lockCollector) defineCollectorAsInitialized() {
p.once.Do(func() {
// mark watcher as initialized.
close(p.initializationC)
})
}
// processEventsAndUpdateCurrent is called when a watcher event is received.
func (p *lockCollector) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
p.currentRW.Lock()
defer p.currentRW.Unlock()
eventsToEmit := events[:0]
for _, event := range events {
if event.Resource == nil || event.Resource.GetKind() != types.KindLock {
p.Logger.WarnContext(ctx, "Received unexpected event", "event", logutils.StringerAttr(event))
continue
}
switch event.Type {
case types.OpDelete:
delete(p.current, event.Resource.GetName())
eventsToEmit = append(eventsToEmit, event)
case types.OpPut:
lock, ok := event.Resource.(types.Lock)
if !ok {
p.Logger.WarnContext(ctx, "Unexpected resource type", "resource", event.Resource.GetKind())
continue
}
if lock.IsInForce(p.Clock.Now()) {
p.current[lock.GetName()] = lock
eventsToEmit = append(eventsToEmit, event)
} else {
delete(p.current, lock.GetName())
}
default:
p.Logger.WarnContext(ctx, "Skipping unsupported event type", "event_type", event.Type)
}
}
p.fanout.Emit(eventsToEmit...)
}
// notifyStale is called when the maximum acceptable staleness (if specified)
// is exceeded.
func (p *lockCollector) notifyStale() {
p.currentRW.Lock()
defer p.currentRW.Unlock()
p.fanout.Emit(types.Event{Type: types.OpUnreliable})
// Do not clear p.current here, the most recent lock set may still be used
// with LockingModeBestEffort.
p.isStale = true
}
func lockTargetsToWatchKinds(targets []types.LockTarget) ([]types.WatchKind, error) {
watchKinds := make([]types.WatchKind, 0, len(targets))
for _, target := range targets {
if target == (types.LockTarget{}) {
continue
}
filter, err := target.IntoMap()
if err != nil {
return nil, trace.Wrap(err)
}
watchKinds = append(watchKinds, types.WatchKind{
Kind: types.KindLock,
Filter: filter,
})
}
if len(watchKinds) == 0 {
watchKinds = []types.WatchKind{{Kind: types.KindLock}}
}
return watchKinds, nil
}
func lockMapValues(lockMap map[string]types.Lock) []types.Lock {
locks := make([]types.Lock, 0, len(lockMap))
for _, lock := range lockMap {
locks = append(locks, lock)
}
return locks
}
func resourcesToSlice[T any](resources map[string]T, cloneFunc func(T) T) (slice []T) {
for _, resource := range resources {
slice = append(slice, cloneFunc(resource))
}
return slice
}
// CertAuthorityWatcherConfig is a CertAuthorityWatcher configuration.
type CertAuthorityWatcherConfig struct {
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
// AuthorityGetter is responsible for fetching cert authority resources.
AuthorityGetter
// Types restricts which cert authority types are retrieved via the AuthorityGetter.
Types []types.CertAuthType
// LoadKeys determines whether private keys will be included.
LoadKeys bool
// ResourceC receives an up-to-date list of all cert authority resources.
ResourceC chan []types.CertAuthority
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *CertAuthorityWatcherConfig) CheckAndSetDefaults() error {
if err := cfg.ResourceWatcherConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if cfg.AuthorityGetter == nil {
getter, ok := cfg.Client.(AuthorityGetter)
if !ok {
return trace.BadParameter("missing parameter AuthorityGetter and Client not usable as AuthorityGetter")
}
cfg.AuthorityGetter = getter
}
if len(cfg.Types) == 0 {
return trace.BadParameter("missing parameter Types")
}
return nil
}
// NewCertAuthorityWatcher returns a new cert authority watcher instance.
func NewCertAuthorityWatcher(ctx context.Context, cfg CertAuthorityWatcherConfig) (*GenericWatcher[types.CertAuthority, readonly.CertAuthority], error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
getter := cfg.AuthorityGetter
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.CertAuthority, readonly.CertAuthority]{
ResourceKind: types.KindCertAuthority,
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceGetter: func(ctx context.Context) ([]types.CertAuthority, error) {
var cas []types.CertAuthority
for _, t := range cfg.Types {
innerCAs, err := getter.GetCertAuthorities(ctx, t, cfg.LoadKeys)
if err != nil {
return nil, trace.Wrap(err)
}
cas = append(cas, innerCAs...)
}
return cas, nil
},
ResourceKey: func(resource types.CertAuthority) string {
return resource.GetSubKind() + "/" + resource.GetName()
},
DeleteKey: func(resource types.Resource) string {
return resource.GetSubKind() + "/" + resource.GetName()
},
ResourcesC: cfg.ResourceC,
CloneFunc: types.CertAuthority.Clone,
ReadOnlyFunc: func(resource types.CertAuthority) readonly.CertAuthority {
return resource
},
LoadSecrets: cfg.LoadKeys,
})
return w, trace.Wrap(err)
}
// DeprecatedNewCertAuthorityWatcher returns a new instance of CertAuthorityWatcher.
//
// Deprecated: This has been replaced by NewCertAuthorityWatcher which uses the
// newer generic watcher under the hood.
func DeprecatedNewCertAuthorityWatcher(ctx context.Context, cfg CertAuthorityWatcherConfig) (*CertAuthorityWatcher, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
collector := &caCollector{
CertAuthorityWatcherConfig: cfg,
fanout: NewFanoutV2(FanoutV2Config{
Capacity: smallFanoutCapacity,
}),
cas: make(map[types.CertAuthType]map[string]types.CertAuthority, len(cfg.Types)),
filter: make(types.CertAuthorityFilter, len(cfg.Types)),
initializationC: make(chan struct{}),
}
for _, t := range cfg.Types {
collector.cas[t] = make(map[string]types.CertAuthority)
collector.filter[t] = types.Wildcard
}
// Resource watcher require the fanout to be initialized before passing in.
// Otherwise, Emit() may fail due to a race condition mentioned in https://github.com/gravitational/teleport/issues/19289
collector.fanout.SetInit(collector.resourceKinds())
watcher, err := newResourceWatcher(ctx, collector, cfg.ResourceWatcherConfig)
if err != nil {
return nil, trace.Wrap(err)
}
return &CertAuthorityWatcher{watcher, collector}, nil
}
// CertAuthorityWatcher is built on top of resourceWatcher to monitor cert authority resources.
type CertAuthorityWatcher struct {
*resourceWatcher
*caCollector
}
// caCollector accompanies resourceWatcher when monitoring cert authority resources.
type caCollector struct {
fanout *FanoutV2
cas map[types.CertAuthType]map[string]types.CertAuthority
// initializationC is used to check whether the initial sync has completed
initializationC chan struct{}
filter types.CertAuthorityFilter
CertAuthorityWatcherConfig
// lock protects concurrent access to cas
lock sync.RWMutex
once sync.Once
}
// Subscribe is used to subscribe to the lock updates.
func (c *caCollector) Subscribe(ctx context.Context, filter types.CertAuthorityFilter) (types.Watcher, error) {
if len(filter) == 0 {
filter = c.filter
}
watch := types.Watch{
Kinds: []types.WatchKind{
{
Kind: types.KindCertAuthority,
Filter: filter.IntoMap(),
},
},
}
sub, err := c.fanout.NewWatcher(ctx, watch)
if err != nil {
return nil, trace.Wrap(err)
}
select {
case event := <-sub.Events():
if event.Type != types.OpInit {
return nil, trace.BadParameter("expected init event, got %v instead", event.Type)
}
case <-sub.Done():
return nil, trace.Wrap(sub.Error())
}
return sub, nil
}
// resourceKinds specifies the resource kind to watch.
func (c *caCollector) resourceKinds() []types.WatchKind {
return []types.WatchKind{{Kind: types.KindCertAuthority, Filter: c.filter.IntoMap()}}
}
// isInitialized is used to check that the cache has done its initial
// sync
func (c *caCollector) initializationChan() <-chan struct{} {
return c.initializationC
}
// getResourcesAndUpdateCurrent refreshes the list of current resources.
func (c *caCollector) getResourcesAndUpdateCurrent(ctx context.Context) error {
var cas []types.CertAuthority
for _, t := range c.Types {
authorities, err := c.AuthorityGetter.GetCertAuthorities(ctx, t, false)
if err != nil {
return trace.Wrap(err)
}
cas = append(cas, authorities...)
}
c.lock.Lock()
defer c.lock.Unlock()
for _, ca := range cas {
if !c.watchingType(ca.GetType()) {
continue
}
c.cas[ca.GetType()][ca.GetName()] = ca
c.fanout.Emit(types.Event{Type: types.OpPut, Resource: ca.Clone()})
}
c.defineCollectorAsInitialized()
return nil
}
func (c *caCollector) defineCollectorAsInitialized() {
c.once.Do(func() {
// mark watcher as initialized.
close(c.initializationC)
})
}
// processEventsAndUpdateCurrent is called when a watcher event is received.
func (c *caCollector) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
c.lock.Lock()
defer c.lock.Unlock()
eventsToEmit := events[:0]
for _, event := range events {
if event.Resource == nil || event.Resource.GetKind() != types.KindCertAuthority {
c.Logger.WarnContext(ctx, "Received unexpected event", "event", logutils.StringerAttr(event))
continue
}
switch event.Type {
case types.OpDelete:
caType := types.CertAuthType(event.Resource.GetSubKind())
if !c.watchingType(caType) {
continue
}
delete(c.cas[caType], event.Resource.GetName())
eventsToEmit = append(eventsToEmit, event)
case types.OpPut:
ca, ok := event.Resource.(types.CertAuthority)
if !ok {
c.Logger.WarnContext(ctx, "Received unexpected resource type", "resource", event.Resource.GetKind())
continue
}
if !c.watchingType(ca.GetType()) {
continue
}
authority, ok := c.cas[ca.GetType()][ca.GetName()]
if ok && authority.IsEqual(ca) {
continue
}
c.cas[ca.GetType()][ca.GetName()] = ca
eventsToEmit = append(eventsToEmit, event)
default:
c.Logger.WarnContext(ctx, "Received unsupported event type", "event_type", event.Type)
}
}
c.fanout.Emit(eventsToEmit...)
}
func (c *caCollector) watchingType(t types.CertAuthType) bool {
if _, ok := c.cas[t]; ok {
return true
}
return false
}
func (c *caCollector) notifyStale() {}
// NodeWatcherConfig is a NodeWatcher configuration.
type NodeWatcherConfig struct {
// NodesGetter is used to directly fetch the list of active nodes.
NodesGetter
ResourceWatcherConfig
// ScopeFilter is an optional scope filter applied to the node watch. See
// GenericWatcherConfig.ScopeFilter for discussion of function. Node watchers
// that run on the teleport proxy need to see nodes in every scope for routing
// and should set this to MODE_ALL.
ScopeFilter *types.ScopeFilter
}
// NewNodeWatcher returns a new instance of NodeWatcher.
func NewNodeWatcher(ctx context.Context, cfg NodeWatcherConfig) (*GenericWatcher[types.Server, readonly.Server], error) {
if cfg.NodesGetter == nil {
return nil, trace.BadParameter("NodesGetter must be provided")
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.Server, readonly.Server]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindNode,
ScopeFilter: cfg.ScopeFilter,
ResourceGetter: func(ctx context.Context) ([]types.Server, error) {
return iterstream.Collect(cfg.NodesGetter.RangeSSHServers(ctx, presencev1.ListSSHServersRequest_builder{
ScopeFilter: cfg.ScopeFilter.ToProto(),
}.Build()))
},
ResourceKey: GetCursorForNode,
DeleteKey: func(res types.Resource) string {
// Delete events for scoped nodes carry a partially populated
// server so that the scope can be recovered from the key, while
// unscoped nodes only emit a resource header.
if node, ok := res.(types.Server); ok {
return GetCursorForNode(node)
}
return scopes.MakeResourceCursor("", res.GetName())
},
DisableUpdateBroadcast: true,
CloneFunc: types.Server.DeepCopy,
ReadOnlyFunc: func(resource types.Server) readonly.Server {
return resource
},
})
return w, trace.Wrap(err)
}
// AccessRequestWatcherConfig is a AccessRequestWatcher configuration.
type AccessRequestWatcherConfig struct {
// AccessRequestGetter is responsible for fetching access request resources.
AccessRequestGetter
// AccessRequestsC receives up-to-date list of all access request resources.
AccessRequestsC chan types.AccessRequests
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig
// Filter is the filter to use to monitor access requests.
Filter types.AccessRequestFilter
}
// CheckAndSetDefaults checks parameters and sets default values.
func (cfg *AccessRequestWatcherConfig) CheckAndSetDefaults() error {
if err := cfg.ResourceWatcherConfig.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if cfg.AccessRequestGetter == nil {
getter, ok := cfg.Client.(AccessRequestGetter)
if !ok {
return trace.BadParameter("missing parameter AccessRequestGetter and Client not usable as AccessRequestGetter")
}
cfg.AccessRequestGetter = getter
}
if cfg.AccessRequestsC == nil {
cfg.AccessRequestsC = make(chan types.AccessRequests)
}
return nil
}
// NewAccessRequestWatcher returns a new instance of AccessRequestWatcher.
func NewAccessRequestWatcher(ctx context.Context, cfg AccessRequestWatcherConfig) (*AccessRequestWatcher, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
collector := &accessRequestCollector{
AccessRequestWatcherConfig: cfg,
initializationC: make(chan struct{}),
}
watcher, err := newResourceWatcher(ctx, collector, cfg.ResourceWatcherConfig)
if err != nil {
return nil, trace.Wrap(err)
}
return &AccessRequestWatcher{watcher, collector}, nil
}
// AccessRequestWatcher is built on top of resourceWatcher to monitor access request resources.
type AccessRequestWatcher struct {
*resourceWatcher
*accessRequestCollector
}
// accessRequestCollector accompanies resourceWatcher when monitoring access request resources.
type accessRequestCollector struct {
// AccessRequestWatcherConfig is the watcher configuration.
AccessRequestWatcherConfig
// current holds a map of the currently known access request resources.
current map[string]types.AccessRequest
// initializationC is used to check that the watcher has been initialized properly.
initializationC chan struct{}
// lock protects the "current" map.
lock sync.RWMutex
once sync.Once
}
// resourceKinds specifies the resource kind to watch.
func (p *accessRequestCollector) resourceKinds() []types.WatchKind {
return []types.WatchKind{{Kind: types.KindAccessRequest}}
}
// isInitialized is used to check that the cache has done its initial
// sync
func (p *accessRequestCollector) initializationChan() <-chan struct{} {
return p.initializationC
}
// getResourcesAndUpdateCurrent refreshes the list of current resources.
func (p *accessRequestCollector) getResourcesAndUpdateCurrent(ctx context.Context) error {
accessRequests, err := p.AccessRequestGetter.GetAccessRequests(ctx, p.Filter)
if err != nil {
return trace.Wrap(err)
}
newCurrent := make(map[string]types.AccessRequest, len(accessRequests))
for _, accessRequest := range accessRequests {
newCurrent[accessRequest.GetName()] = accessRequest.Copy()
}
p.lock.Lock()
defer p.lock.Unlock()
p.current = newCurrent
p.defineCollectorAsInitialized()
select {
case <-ctx.Done():
return trace.Wrap(ctx.Err())
case p.AccessRequestsC <- accessRequests:
}
return nil
}
func (p *accessRequestCollector) defineCollectorAsInitialized() {
p.once.Do(func() {
// mark watcher as initialized.
close(p.initializationC)
})
}
// processEventsAndUpdateCurrent is called when a watcher event is received.
func (p *accessRequestCollector) processEventsAndUpdateCurrent(ctx context.Context, events []types.Event) {
p.lock.Lock()
defer p.lock.Unlock()
for _, event := range events {
if event.Resource == nil || event.Resource.GetKind() != types.KindAccessRequest {
p.Logger.WarnContext(ctx, "Received unexpected event", "event", logutils.StringerAttr(event))
continue
}
switch event.Type {
case types.OpDelete:
delete(p.current, event.Resource.GetName())
select {
case <-ctx.Done():
case p.AccessRequestsC <- resourcesToSlice(p.current, types.AccessRequest.Copy):
}
case types.OpPut:
accessRequest, ok := event.Resource.(types.AccessRequest)
if !ok {
p.Logger.WarnContext(ctx, "Received unexpected resource type", "resource", event.Resource.GetKind())
continue
}
p.current[accessRequest.GetName()] = accessRequest
select {
case <-ctx.Done():
case p.AccessRequestsC <- resourcesToSlice(p.current, types.AccessRequest.Copy):
}
default:
p.Logger.WarnContext(ctx, "Received unsupported event type", "event_type", event.Type)
}
}
}
func (*accessRequestCollector) notifyStale() {}
// GitServerWatcherConfig is the config for Git server watcher.
type GitServerWatcherConfig struct {
GitServerGetter
ResourceWatcherConfig
// EnableUpdateBroadcast turns on emitting updates on changes. Broadcast is
// opt-in for Git Server watcher.
EnableUpdateBroadcast bool
}
// NewGitServerWatcher returns a new instance of Git server watcher.
func NewGitServerWatcher(ctx context.Context, cfg GitServerWatcherConfig) (*GenericWatcher[types.Server, readonly.Server], error) {
if cfg.GitServerGetter == nil {
return nil, trace.BadParameter("GitServerGetter must be provided")
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[types.Server, readonly.Server]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindGitServer,
ResourceGetter: pagerFn[types.Server](
cfg.GitServerGetter.ListGitServers,
).getAll,
ResourceKey: types.Server.GetName,
DisableUpdateBroadcast: !cfg.EnableUpdateBroadcast,
CloneFunc: types.Server.DeepCopy,
ReadOnlyFunc: func(resource types.Server) readonly.Server {
return resource
},
})
return w, trace.Wrap(err)
}
// HealthCheckConfigWatcherConfig is the config for the health_check_config
// watcher.
type HealthCheckConfigWatcherConfig struct {
// Reader is used to fetch health check config resources.
Reader HealthCheckConfigReader
// ResourcesC receives up-to-date list of all health check config resources.
ResourcesC chan []*healthcheckconfigv1.HealthCheckConfig
// ResourceWatcherConfig is the resource watcher configuration.
ResourceWatcherConfig ResourceWatcherConfig
}
// HealthCheckConfigWatcher monitors health_check_config resources.
type HealthCheckConfigWatcher = GenericWatcher[
*healthcheckconfigv1.HealthCheckConfig,
*healthcheckconfigv1.HealthCheckConfig,
]
// NewHealthCheckConfigWatcher returns a new instance of health check config
// watcher.
func NewHealthCheckConfigWatcher(
ctx context.Context,
cfg HealthCheckConfigWatcherConfig,
) (*HealthCheckConfigWatcher, error) {
if cfg.Reader == nil {
return nil, trace.BadParameter("Reader must be provided")
}
w, err := NewGenericResourceWatcher(ctx, GenericWatcherConfig[
*healthcheckconfigv1.HealthCheckConfig,
*healthcheckconfigv1.HealthCheckConfig,
]{
ResourceWatcherConfig: cfg.ResourceWatcherConfig,
ResourceKind: types.KindHealthCheckConfig,
ResourceGetter: pagerFn[*healthcheckconfigv1.HealthCheckConfig](
cfg.Reader.ListHealthCheckConfigs,
).getAll,
ResourceKey: func(resource *healthcheckconfigv1.HealthCheckConfig) string {
return resource.GetMetadata().GetName()
},
ResourcesC: cfg.ResourcesC,
CloneFunc: apiutils.CloneProtoMsg[*healthcheckconfigv1.HealthCheckConfig],
ReadOnlyFunc: func(resource *healthcheckconfigv1.HealthCheckConfig) *healthcheckconfigv1.HealthCheckConfig {
return resource
},
})
return w, trace.Wrap(err)
}
type pagerFn[T any] func(ctx context.Context, limit int, startKey string) ([]T, string, error)
func (fn pagerFn[T]) getAll(ctx context.Context) ([]T, error) {
var out []T
var token string
for {
page, nextToken, err := fn(ctx, apidefaults.DefaultChunkSize, token)
if err != nil {
return nil, trace.Wrap(err)
}
out = append(out, page...)
if nextToken == "" {
return out, nil
}
token = nextToken
}
}
// Teleport
// Copyright (C) 2024 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"
"encoding/base32"
"iter"
"slices"
"strings"
"time"
"github.com/gravitational/trace"
"golang.org/x/text/cases"
workloadidentityv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/auth/machineid/workloadidentityv1/expression"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/scopes"
)
// WorkloadIdentities is an interface over the WorkloadIdentities service. This
// interface may also be implemented by a client to allow remote and local
// consumers to access the resource in a similar way.
type WorkloadIdentities interface {
// GetWorkloadIdentity gets a WorkloadIdentity by the name and scope in the
// request. An empty scope addresses an unscoped WorkloadIdentity.
GetWorkloadIdentity(
ctx context.Context, req *workloadidentityv1pb.GetWorkloadIdentityRequest,
) (*workloadidentityv1pb.WorkloadIdentity, error)
// RangeWorkloadIdentities returns WorkloadIdentity resources within the
// range [start, end), ordered by the given sort field and direction.
RangeWorkloadIdentities(
ctx context.Context,
start, end string,
sortField WorkloadIdentitySortField,
sortDesc bool,
) iter.Seq2[*workloadidentityv1pb.WorkloadIdentity, error]
// CreateWorkloadIdentity creates a new WorkloadIdentity.
CreateWorkloadIdentity(
ctx context.Context, workloadIdentity *workloadidentityv1pb.WorkloadIdentity,
) (*workloadidentityv1pb.WorkloadIdentity, error)
// DeleteWorkloadIdentity deletes a WorkloadIdentity by the name and scope in
// the request.
DeleteWorkloadIdentity(ctx context.Context, req *workloadidentityv1pb.DeleteWorkloadIdentityRequest) error
// UpdateWorkloadIdentity updates a specific WorkloadIdentity. The resource must
// already exist, and, condition update semantics are used - e.g the submitted
// resource must have a revision matching the revision of the resource in the
// backend.
UpdateWorkloadIdentity(
ctx context.Context, workloadIdentity *workloadidentityv1pb.WorkloadIdentity,
) (*workloadidentityv1pb.WorkloadIdentity, error)
// UpsertWorkloadIdentity creates or updates a WorkloadIdentity.
UpsertWorkloadIdentity(
ctx context.Context, workloadIdentity *workloadidentityv1pb.WorkloadIdentity,
) (*workloadidentityv1pb.WorkloadIdentity, error)
// AppendPutWorkloadIdentityActions adds conditional actions to an atomic
// write to create or update a WorkloadIdentity.
AppendPutWorkloadIdentityActions(
actions []backend.ConditionalAction,
resource *workloadidentityv1pb.WorkloadIdentity,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
// AppendDeleteWorkloadIdentityActions adds conditional actions to an atomic
// write to delete a WorkloadIdentity given its scope-qualified name.
AppendDeleteWorkloadIdentityActions(
actions []backend.ConditionalAction,
name scopes.QualifiedName,
condition backend.Condition,
) ([]backend.ConditionalAction, error)
}
// MarshalWorkloadIdentity marshals the WorkloadIdentity object into a JSON byte
// array.
func MarshalWorkloadIdentity(
object *workloadidentityv1pb.WorkloadIdentity, opts ...MarshalOption,
) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalWorkloadIdentity unmarshals the WorkloadIdentity object from a
// JSON byte array.
func UnmarshalWorkloadIdentity(
data []byte, opts ...MarshalOption,
) (*workloadidentityv1pb.WorkloadIdentity, error) {
return UnmarshalProtoResource[*workloadidentityv1pb.WorkloadIdentity](data, opts...)
}
const (
maxMaxJWTSVIDTTL = time.Hour * 24
maxMaxX509SVIDTTL = time.Hour * 24 * 14
)
// ValidateWorkloadIdentity validates the WorkloadIdentity object. This is
// performed prior to writing to the backend.
func ValidateWorkloadIdentity(s *workloadidentityv1pb.WorkloadIdentity) error {
switch {
case s == nil:
return trace.BadParameter("object cannot be nil")
case s.GetVersion() != types.V1:
return trace.BadParameter("version: only %q is supported", types.V1)
case s.GetKind() != types.KindWorkloadIdentity:
return trace.BadParameter("kind: must be %q", types.KindWorkloadIdentity)
case !s.HasMetadata():
return trace.BadParameter("metadata: is required")
case s.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name: is required")
case !s.HasSpec():
return trace.BadParameter("spec: is required")
case s.GetSpec().GetSpiffe().GetId() == "":
return trace.BadParameter("spec.spiffe.id: is required")
case !strings.HasPrefix(s.GetSpec().GetSpiffe().GetId(), "/"):
return trace.BadParameter("spec.spiffe.id: must start with a /")
case s.GetSpec().GetSpiffe().GetX509().GetMaximumTtl().AsDuration() > maxMaxX509SVIDTTL:
return trace.BadParameter("spec.spiffe.x509.maximum_ttl: must be less than %s", maxMaxX509SVIDTTL)
case s.GetSpec().GetSpiffe().GetJwt().GetMaximumTtl().AsDuration() > maxMaxJWTSVIDTTL:
return trace.BadParameter("spec.spiffe.jwt.maximum_ttl: must be less than %s", maxMaxJWTSVIDTTL)
}
// When the WorkloadIdentity is scoped, the scope itself must be valid and
// the SPIFFE ID must conform to the scoped SPIFFE ID structure (RFD 0229c).
if s.GetScope() != "" {
if err := scopes.StrongValidate(s.GetScope()); err != nil {
return trace.Wrap(err, "scope")
}
if scopes.Compare(s.GetScope(), scopes.Root) == scopes.Equivalent {
return trace.BadParameter("scope: must not be the root scope")
}
if err := ValidateScopedSPIFFEID(s.GetScope(), s.GetSpec().GetSpiffe().GetId()); err != nil {
return trace.Wrap(err)
}
// TODO(strideynet): For now we only constrict the naming of scoped
// workload identities - however - we should consider rolling out a
// write-side restriction to unscoped workload identities in a major
// version.
if err := scopes.StrongValidateResourceName(s.GetMetadata().GetName()); err != nil {
return trace.Wrap(err, "metadata.name:")
}
}
for i, rule := range s.GetSpec().GetRules().GetAllow() {
if rule.GetExpression() == "" {
if len(rule.GetConditions()) == 0 {
return trace.BadParameter("spec.rules.allow[%d].conditions: must be non-empty", i)
}
} else {
if len(rule.GetConditions()) != 0 {
return trace.BadParameter("spec.rules.allow[%d].conditions: is mutually exclusive with expression", i)
}
if err := expression.Validate(rule.GetExpression()); err != nil {
return trace.BadParameter("spec.rules.allow[%d].expression: invalid expression: %s", i, err.Error())
}
}
for j, condition := range rule.GetConditions() {
if condition.GetAttribute() == "" {
return trace.BadParameter("spec.rules.allow[%d].conditions[%d].attribute: must be non-empty", i, j)
}
if !condition.HasOperator() {
return trace.BadParameter("spec.rules.allow[%d].conditions[%d]: operator must be specified", i, j)
}
}
}
return nil
}
// scopedSPIFFEIDSeparator is the path segment that separates the scope-derived
// section of a scoped SPIFFE ID from the administratively-defined section. A
// scoped SPIFFE ID for a WorkloadIdentity defined in scope /security/eu looks
// like /security/eu/_/k8s/cluster-a. See RFD 0229c.
//
// The separator is a valid SPIFFE ID path segment but is deliberately not a
// valid scope segment (scope segments require at least two characters), which
// keeps the boundary between the two sections unambiguous.
const scopedSPIFFEIDSeparator = "_"
// ValidateScopedSPIFFEID validates that the given SPIFFE ID path conforms to
// the scoped SPIFFE ID structure for the given scope, as defined in RFD 0229c.
//
// A scoped SPIFFE ID path consists of three sections:
// - the scope section: segments that strictly match the scope of origin
// - the separator segment ("/_/")
// - the administratively-defined section: one or more freely-defined segments
//
// For example, for scope /security/eu, /security/eu/_/k8s/cluster-a is valid.
//
// The check is performed on segments (not raw string prefixes) so that, e.g., a
// scope of /foo matches an ID beginning /foo/... but not /foo-buzz/.... The
// scope section must match the scope of origin exactly: it may not be an
// ancestor or descendant of it.
//
// It is used both at create/update time (on the unrendered ID) and at issuance
// time (on the rendered ID) as defense-in-depth against templating that would
// otherwise escape the WorkloadIdentity's scope.
func ValidateScopedSPIFFEID(scope, id string) error {
if !strings.HasPrefix(id, "/") {
return trace.BadParameter("spec.spiffe.id: must begin with a forward slash")
}
// Reject empty path segments (e.g. a trailing slash or "//"), which
// splitPathSegments would otherwise silently trim or mishandle.
if strings.HasSuffix(id, "/") || strings.Contains(id, "//") {
return trace.BadParameter("spec.spiffe.id %q must not contain empty path segments", id)
}
scopeSegments := splitPathSegments(scope)
idSegments := splitPathSegments(id)
separatorIndex := slices.Index(idSegments, scopedSPIFFEIDSeparator)
if separatorIndex < 0 {
return trace.BadParameter(
"spec.spiffe.id %q is missing the %q separator segment that delimits the scope from the administratively-defined section",
id, scopedSPIFFEIDSeparator,
)
}
scopeSection := idSegments[:separatorIndex]
adminSection := idSegments[separatorIndex+1:]
if !slices.Equal(scopeSection, scopeSegments) {
return trace.BadParameter(
"spec.spiffe.id %q must be prefixed with the scope %q, immediately followed by the %q separator segment",
id, scope, scopedSPIFFEIDSeparator,
)
}
if len(adminSection) == 0 {
return trace.BadParameter(
"spec.spiffe.id %q must have at least one segment after the %q separator",
id, scopedSPIFFEIDSeparator,
)
}
if slices.Contains(adminSection, scopedSPIFFEIDSeparator) {
return trace.BadParameter(
"spec.spiffe.id %q must not contain the %q separator segment in its administratively-defined section",
id, scopedSPIFFEIDSeparator,
)
}
return nil
}
// splitPathSegments splits a slash-delimited path (a scope or a SPIFFE ID path)
// into its non-empty segments. A leading slash is expected; empty input or "/"
// yields no segments.
func splitPathSegments(path string) []string {
trimmed := strings.Trim(path, "/")
if trimmed == "" {
return nil
}
return strings.Split(trimmed, "/")
}
// WorkloadIdentitySortField identifies a field that WorkloadIdentities may be
// sorted (and ranged) by. An empty value defaults to
// [WorkloadIdentitySortFieldName].
type WorkloadIdentitySortField string
const (
// WorkloadIdentitySortFieldName sorts WorkloadIdentities by name. It is the
// default when no sort field is specified.
WorkloadIdentitySortFieldName WorkloadIdentitySortField = "name"
// WorkloadIdentitySortFieldSPIFFEID sorts WorkloadIdentities by SPIFFE ID.
WorkloadIdentitySortFieldSPIFFEID WorkloadIdentitySortField = "spiffe_id"
)
// WorkloadIdentityKey returns a function deriving the canonical ordering key for
// a WorkloadIdentity for the given sort field. It is the single source of truth
// for both the order in which WorkloadIdentities are iterated and the pagination
// cursor used to resume iteration. Supported sort fields are "" (defaults to
// name), [WorkloadIdentitySortFieldName] and [WorkloadIdentitySortFieldSPIFFEID];
// any other sort field returns an error.
//
// The sort field is validated once, when the key function is obtained, so
// callers do not have to handle an error per resource.
func WorkloadIdentityKey(sortField WorkloadIdentitySortField) (func(*workloadidentityv1pb.WorkloadIdentity) string, error) {
switch sortField {
case "", WorkloadIdentitySortFieldName:
return workloadIdentityCursor, nil
case WorkloadIdentitySortFieldSPIFFEID:
return workloadIdentitySPIFFEIDKey, nil
default:
return nil, trace.BadParameter("unsupported sort %q but expected %s or %s", sortField, WorkloadIdentitySortFieldName, WorkloadIdentitySortFieldSPIFFEID)
}
}
// workloadIdentityCursor returns the canonical resource cursor for a
// WorkloadIdentity, used both as its in-memory cache index key and as its
// pagination cursor, so it must be stable and unique per resource.
func workloadIdentityCursor(wi *workloadidentityv1pb.WorkloadIdentity) string {
return scopes.MakeResourceCursor(wi.GetScope(), wi.GetMetadata().GetName())
}
// workloadIdentitySPIFFEIDKey returns the ordering key for the spiffe_id sort.
func workloadIdentitySPIFFEIDKey(wi *workloadidentityv1pb.WorkloadIdentity) string {
// Sort case-insensitively to keep /spiffe-1 and /Spiffe-1 together.
spiffeID := cases.Fold().String(wi.GetSpec().GetSpiffe().GetId())
// Encode to avoid ambiguity; "a/b" + "/" + "c" vs. "a" + "/" + "b/c". Base32
// hex maintains the original ordering.
spiffeID = base32.HexEncoding.WithPadding(base32.NoPadding).EncodeToString([]byte(spiffeID))
// SPIFFE IDs may not be unique, so append the resource cursor, which
// uniquely identifies the resource across scopes.
return spiffeID + "/" + workloadIdentityCursor(wi)
}
// Teleport
// Copyright (C) 2025 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"
"regexp"
"github.com/gravitational/trace"
workloadidentityv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/api/types"
)
// WorkloadIdentityX509Revocations is an interface over the
// WorkloadIdentityX509Revocations service. This interface may also be
// implemented by a client to allow remote and local consumers to access the
// resource in a similar way.
type WorkloadIdentityX509Revocations interface {
// GetWorkloadIdentityX509Revocation gets a WorkloadIdentityX509Revocation
// by name.
GetWorkloadIdentityX509Revocation(
ctx context.Context, name string,
) (*workloadidentityv1pb.WorkloadIdentityX509Revocation, error)
// ListWorkloadIdentityX509Revocations lists all
// WorkloadIdentityX509Revocation using Google style pagination.
ListWorkloadIdentityX509Revocations(
ctx context.Context, pageSize int, lastToken string,
) ([]*workloadidentityv1pb.WorkloadIdentityX509Revocation, string, error)
// CreateWorkloadIdentityX509Revocation creates a new
// WorkloadIdentityX509Revocation.
CreateWorkloadIdentityX509Revocation(
ctx context.Context,
workloadIdentityX509Revocation *workloadidentityv1pb.WorkloadIdentityX509Revocation,
) (*workloadidentityv1pb.WorkloadIdentityX509Revocation, error)
// DeleteWorkloadIdentityX509Revocation deletes a
// WorkloadIdentityX509Revocation by name.
DeleteWorkloadIdentityX509Revocation(ctx context.Context, name string) error
// UpdateWorkloadIdentityX509Revocation updates a specific
// WorkloadIdentityX509Revocation. The resource must already exist, and,
// conditional update semantics are used - e.g the submitted resource must
// have a revision matching the revision of the resource in the backend.
UpdateWorkloadIdentityX509Revocation(
ctx context.Context,
workloadIdentityX509Revocation *workloadidentityv1pb.WorkloadIdentityX509Revocation,
) (*workloadidentityv1pb.WorkloadIdentityX509Revocation, error)
// UpsertWorkloadIdentityX509Revocation creates or updates a
// WorkloadIdentityX509Revocation.
UpsertWorkloadIdentityX509Revocation(
ctx context.Context,
workloadIdentityX509Revocation *workloadidentityv1pb.WorkloadIdentityX509Revocation,
) (*workloadidentityv1pb.WorkloadIdentityX509Revocation, error)
}
// MarshalWorkloadIdentityX509Revocation marshals the
// WorkloadIdentityX509Revocation object into a JSON byte array.
func MarshalWorkloadIdentityX509Revocation(
object *workloadidentityv1pb.WorkloadIdentityX509Revocation, opts ...MarshalOption,
) ([]byte, error) {
return MarshalProtoResource(object, opts...)
}
// UnmarshalWorkloadIdentityX509Revocation unmarshals the
// WorkloadIdentityX509Revocation object from a JSON byte array.
func UnmarshalWorkloadIdentityX509Revocation(
data []byte, opts ...MarshalOption,
) (*workloadidentityv1pb.WorkloadIdentityX509Revocation, error) {
return UnmarshalProtoResource[*workloadidentityv1pb.WorkloadIdentityX509Revocation](data, opts...)
}
var validSerialRe = regexp.MustCompile("^(?:[0-9a-f]{2})+$")
// ValidateWorkloadIdentityX509Revocation validates the
// WorkloadIdentityX509Revocation object.
// It returns a nil if the object is valid, otherwise an error.
func ValidateWorkloadIdentityX509Revocation(s *workloadidentityv1pb.WorkloadIdentityX509Revocation) error {
switch {
case s == nil:
return trace.BadParameter("object cannot be nil")
case s.GetVersion() != types.V1:
return trace.BadParameter("version: only %q is supported", types.V1)
case s.GetKind() != types.KindWorkloadIdentityX509Revocation:
return trace.BadParameter("kind: must be %q", types.KindWorkloadIdentityX509Revocation)
case !s.HasMetadata():
return trace.BadParameter("metadata: is required")
case s.GetMetadata().GetName() == "":
return trace.BadParameter("metadata.name: is required")
case !s.GetMetadata().HasExpires():
return trace.BadParameter("metadata.expires: is required")
case s.GetMetadata().GetExpires().IsValid() == false:
return trace.BadParameter("metadata.expires: must be valid")
case s.GetMetadata().GetExpires().AsTime().IsZero():
return trace.BadParameter("metadata.expires: must be non-zero")
case !s.HasSpec():
return trace.BadParameter("spec: is required")
case s.GetSpec().GetReason() == "":
return trace.BadParameter("spec.reason: is required")
case !s.GetSpec().HasRevokedAt():
return trace.BadParameter("spec.revoked_at: is required")
case s.GetSpec().GetRevokedAt().IsValid() == false:
return trace.BadParameter("spec.revoked_at: must be valid")
case s.GetSpec().GetRevokedAt().AsTime().IsZero():
return trace.BadParameter("spec.revoked_at: must be non-zero")
}
// Name must be a integer encoded as hex - this is the serial number of the
// X509 cert. Whilst typically presented using a colon separated hex string,
// here we will remove the colons. We will also ensure it is encoded in
// lowercase, to ensure consistency.
if !validSerialRe.MatchString(s.GetMetadata().GetName()) {
return trace.BadParameter("metadata.name: must be a lower-case hex encoded integer without colons")
}
return nil
}
/*
* 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 protocol
import (
"bytes"
"io"
"net"
"sync"
"github.com/datastax/go-cassandra-native-protocol/client"
"github.com/datastax/go-cassandra-native-protocol/frame"
"github.com/datastax/go-cassandra-native-protocol/message"
"github.com/datastax/go-cassandra-native-protocol/primitive"
"github.com/datastax/go-cassandra-native-protocol/segment"
"github.com/gravitational/trace"
)
// NewConn is used to create a new connection.
func NewConn(rawConn net.Conn) *Conn {
return &Conn{
Conn: rawConn,
frameCodec: frame.NewRawCodec(),
segmentCode: segment.NewCodec(),
}
}
// Conn represent incoming client connection or ongoing connection to cassandra server.
// Reading and Writing Frames/Packages needs to be done sequential because Conn implementation
// is not thread safe. Conn package is used to intercept and preform custom Cassandra handshake
// and in case of connection incoming connection to provide ability to audit incoming client packages.
type Conn struct {
net.Conn
segmentCode segment.Codec
frameCodec frame.RawCodec
modernLayoutRead bool
modernLayoutWrite bool
// selfContained is used to store frames that are self-contained.
// allowing to read them before reading from the connection.
selfContained []*frame.Frame
mtx sync.Mutex
}
// ReadPacket is used to read packet from the connection.
func (c *Conn) ReadPacket() (*Packet, error) {
if fr := c.checkSelfContained(); fr != nil {
return &Packet{frame: fr}, nil
}
var buff bytes.Buffer
tr := io.TeeReader(c.Conn, &buff)
fr, err := c.readWithModernLayout(tr)
if err != nil {
return nil, trace.Wrap(err)
}
c.maybeSwitchToModernLayout(fr)
return &Packet{
raw: buff,
frame: fr,
}, nil
}
// WriteFrame is used to write frame to the connection.
func (c *Conn) WriteFrame(outgoing *frame.Frame) error {
if err := c.writeFrameWithModernLayout(outgoing); err != nil {
return trace.Wrap(err)
}
if startup, ok := outgoing.Body.Message.(*message.Startup); ok {
compression := startup.GetCompression()
c.frameCodec = frame.NewRawCodecWithCompression(client.NewBodyCompressor(compression))
c.segmentCode = segment.NewCodecWithCompression(client.NewPayloadCompressor(compression))
}
c.maybeSwitchToModernLayout(outgoing)
return nil
}
// readWithModernLayout is used to read frame from the connection.
// If the connection is using modern framing layout, it will read segments.
// Otherwise, it will read frames.
func (c *Conn) readWithModernLayout(r io.Reader) (*frame.Frame, error) {
if c.modernLayoutRead {
v, err := c.readSegment(r)
return v, trace.Wrap(err)
}
fr, err := c.readFrame(r)
return fr, trace.Wrap(err)
}
// readFrame is used to read frame from the connection.
// If read frame is a Startup frame, it will switch to modern framing layout and
// update codec to use modern framing layout.
func (c *Conn) readFrame(r io.Reader) (*frame.Frame, error) {
fr, err := c.frameCodec.DecodeFrame(r)
if err != nil {
return nil, trace.Wrap(err)
}
if startup, ok := fr.Body.Message.(*message.Startup); ok {
compression := startup.GetCompression()
if !compression.IsValid() {
return nil, trace.BadParameter("invalid compression: %v", compression)
}
c.frameCodec = frame.NewRawCodecWithCompression(client.NewBodyCompressor(compression))
c.segmentCode = segment.NewCodecWithCompression(client.NewPayloadCompressor(compression))
// If moderate framing layout is supported all received from a client after Startup message should
// use segment encoding.
c.modernLayoutRead = fr.Header.Version.SupportsModernFramingLayout()
}
return fr, nil
}
func (c *Conn) checkSelfContained() *frame.Frame {
c.mtx.Lock()
defer c.mtx.Unlock()
if len(c.selfContained) == 0 {
return nil
}
out := c.selfContained[0]
c.selfContained = c.selfContained[1:]
return out
}
// readSegment is used to read segments from the connection.
// If frame is not self-contained, segments are split into multiple frames.
// Read the segments till received bodyLength bytes.
func (c *Conn) readSegment(r io.Reader) (*frame.Frame, error) {
previousSegment := bytes.Buffer{}
expectedSegmentSize := 0
for {
seg, err := c.segmentCode.DecodeSegment(r)
if err != nil {
return nil, trace.Wrap(err)
}
if seg.Header.IsSelfContained {
fr, err := c.readSelfContainedSegment(seg)
if err != nil {
return nil, trace.Wrap(err)
}
return fr, nil
}
// Otherwise read the frame size and keep reading until we read all data.
// Segments are always delivered in order.
if expectedSegmentSize == 0 {
frameHeader, err := c.frameCodec.DecodeHeader(bytes.NewReader(seg.Payload.UncompressedData))
if err != nil {
return nil, trace.Wrap(err)
}
expectedSegmentSize = int(primitive.FrameHeaderLengthV3AndHigher + frameHeader.BodyLength)
}
// Append another segment
if _, err := previousSegment.Write(seg.Payload.UncompressedData); err != nil {
return nil, trace.Wrap(err)
}
// Return the frame after reading all segments.
if expectedSegmentSize == previousSegment.Len() {
fr, err := c.readFrame(&previousSegment)
if err != nil {
return nil, trace.Wrap(err)
}
return fr, nil
}
}
}
func (c *Conn) readSelfContainedSegment(incoming *segment.Segment) (*frame.Frame, error) {
var out []*frame.Frame
payloadReader := bytes.NewReader(incoming.Payload.UncompressedData)
for payloadReader.Len() > 0 {
fr, err := c.readFrame(payloadReader)
if err != nil {
return nil, trace.Wrap(err)
}
out = append(out, fr)
}
if len(out) == 0 {
return nil, trace.BadParameter("no frames in self-contained segment")
}
c.mtx.Lock()
defer c.mtx.Unlock()
c.selfContained = append(c.selfContained, out[1:]...)
return out[0], nil
}
// writeFrameWithModernLayout is used to write frame to the connection.
func (c *Conn) writeFrameWithModernLayout(outgoing *frame.Frame) error {
if c.modernLayoutWrite {
return trace.Wrap(c.writeSegment(outgoing, c.Conn))
}
return trace.Wrap(c.writeFrame(outgoing, c.Conn))
}
// writeFrame is used to write frame to the connection.
func (c *Conn) writeFrame(outgoing *frame.Frame, wr io.Writer) error {
err := c.frameCodec.EncodeFrame(outgoing, wr)
return trace.Wrap(err)
}
// writeSegment is used to write segments to the connection.
func (c *Conn) writeSegment(outgoing *frame.Frame, wr io.Writer) error {
var buff bytes.Buffer
if err := c.writeFrame(outgoing, &buff); err != nil {
return trace.Wrap(err)
}
seg := &segment.Segment{
Header: &segment.Header{IsSelfContained: true},
Payload: &segment.Payload{UncompressedData: buff.Bytes()},
}
if err := c.segmentCode.EncodeSegment(seg, wr); err != nil {
return trace.Wrap(err)
}
return nil
}
// maybeSwitchToModernLayout is used to switch to modern framing layout.
// If received frame is a Ready frame or Authenticate frame, it will switch to modern framing layout.
func (c *Conn) maybeSwitchToModernLayout(fr *frame.Frame) {
if !isReady(fr) && !isAuthenticate(fr) {
return
}
if !c.modernLayoutRead {
c.modernLayoutRead = fr.Header.Version.SupportsModernFramingLayout()
}
if !c.modernLayoutWrite {
c.modernLayoutWrite = fr.Header.Version.SupportsModernFramingLayout()
}
}
func isReady(fr *frame.Frame) bool {
return fr.Header.OpCode == primitive.OpCodeReady
}
func isAuthenticate(fr *frame.Frame) bool {
return fr.Header.OpCode == primitive.OpCodeAuthenticate
}
/*
* 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 protocol
import (
"bytes"
"github.com/datastax/go-cassandra-native-protocol/frame"
)
// Packet represent cassandra packet frame with
// raw unparsed packet payload.
type Packet struct {
// raw is raw packet payload.
raw bytes.Buffer
// frame is cassandra protocol frame.
frame *frame.Frame
}
// Frame returns frame.
func (p *Packet) Frame() *frame.Frame {
return p.frame
}
// FrameBody returns frame body.
func (p *Packet) FrameBody() *frame.Body {
return p.frame.Body
}
// Header returns frame header.
func (p *Packet) Header() *frame.Header {
return p.frame.Header
}
// Raw returns raw packet payload.
func (p *Packet) Raw() []byte {
return p.raw.Bytes()
}
/*
* 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 elasticsearch
import (
"strings"
apievents "github.com/gravitational/teleport/api/types/events"
)
// parsePath returns (optional) target of query as well as the event category.
func parsePath(path string) (string, apievents.ElasticsearchCategory) {
parts := strings.Split(path, "/")
if len(parts) < 2 {
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_GENERAL
}
// first term starts with _
switch parts[1] {
case "_security", "_ssl":
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_SECURITY
case
"_search", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-search.html
"_async_search", // https://www.elastic.co/guide/en/elasticsearch/reference/master/async-search.html
"_pit", // https://www.elastic.co/guide/en/elasticsearch/reference/master/point-in-time-api.html
"_msearch", // https://www.elastic.co/guide/en/elasticsearch/reference/master/multi-search-template.html, https://www.elastic.co/guide/en/elasticsearch/reference/master/search-multi-search.html
"_render", // https://www.elastic.co/guide/en/elasticsearch/reference/master/render-search-template-api.html
"_field_caps": // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-field-caps.html
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_SEARCH
case "_sql":
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_SQL
}
// starts with _, but we don't handle it explicitly
if strings.HasPrefix("_", parts[1]) {
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_GENERAL
}
if len(parts) < 3 {
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_GENERAL
}
// a number of APIs are invoked by providing a target first, e.g. /<target>/_search, where <target> is an index or expression matching a group of indices.
switch parts[2] {
case
"_search", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-search.html
"_async_search", // https://www.elastic.co/guide/en/elasticsearch/reference/master/async-search.html
"_pit", // https://www.elastic.co/guide/en/elasticsearch/reference/master/point-in-time-api.html
"_knn_search", // https://www.elastic.co/guide/en/elasticsearch/reference/master/knn-search-api.html
"_msearch", // https://www.elastic.co/guide/en/elasticsearch/reference/master/multi-search-template.html, https://www.elastic.co/guide/en/elasticsearch/reference/master/search-multi-search.html
"_search_shards", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-shards.html
"_count", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-count.html
"_validate", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-validate.html
"_terms_enum", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-terms-enum.html
"_explain", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-explain.html
"_field_caps", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-field-caps.html
"_rank_eval", // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-rank-eval.html
"_mvt": // https://www.elastic.co/guide/en/elasticsearch/reference/master/search-vector-tile-api.html
return parts[1], apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_SEARCH
}
return "", apievents.ElasticsearchCategory_ELASTICSEARCH_CATEGORY_GENERAL
}
/*
* 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 elasticsearch
import (
"bufio"
"bytes"
"context"
"encoding/json"
"io"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"github.com/gravitational/trace"
"github.com/prometheus/client_golang/prometheus"
"github.com/gravitational/teleport"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/api/types/wrappers"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/srv/db/common"
"github.com/gravitational/teleport/lib/srv/db/common/role"
"github.com/gravitational/teleport/lib/utils"
)
// NewEngine create new elasticsearch engine.
func NewEngine(ec common.EngineConfig) common.Engine {
return &Engine{EngineConfig: ec}
}
// Engine handles connections from Elasticsearch clients coming from Teleport
// proxy over reverse tunnel.
type Engine struct {
// EngineConfig is the common database engine configuration.
common.EngineConfig
// clientConn is a client connection.
clientConn net.Conn
// sessionCtx is current session context.
sessionCtx *common.Session
}
func (e *Engine) InitializeConnection(clientConn net.Conn, sessionCtx *common.Session) error {
e.clientConn = clientConn
e.sessionCtx = sessionCtx
return nil
}
// SendError sends an error to Elasticsearch client.
func (e *Engine) SendError(err error) {
if e.clientConn == nil || err == nil || utils.IsOKNetworkError(err) {
return
}
// ErrorCause type.
//
// https://github.com/elastic/elasticsearch-specification/blob/f6a370d0fba975752c644fc730f7c45610e28f36/specification/_types/Errors.ts#L25-L50
type ErrorCause struct {
// Reason A human-readable explanation of the error, in English.
Reason *string `json:"reason,omitempty"`
// Type The type of error
Type string `json:"type"`
}
reason := err.Error()
cause := ErrorCause{
Reason: &reason,
Type: "internal_server_error_exception",
}
// Assume internal server error HTTP 500 and override if possible.
statusCode := http.StatusInternalServerError
if trace.IsAccessDenied(err) {
statusCode = http.StatusUnauthorized
cause.Type = "access_denied_exception"
}
jsonBody, err := json.Marshal(cause)
if err != nil {
e.Log.ErrorContext(e.Context, "failed to marshal error response", "error", err)
return
}
response := &http.Response{
ProtoMajor: 1,
ProtoMinor: 1,
StatusCode: statusCode,
Body: io.NopCloser(bytes.NewBuffer(jsonBody)),
Header: map[string][]string{
"Content-Type": {"application/json"},
"Content-Length": {strconv.Itoa(len(jsonBody))},
},
}
if err := response.Write(e.clientConn); err != nil {
e.Log.ErrorContext(e.Context, "elasticsearch error", "error", err)
return
}
}
// HandleConnection authorizes the incoming client connection, connects to the
// target Elasticsearch server and starts proxying requests between client/server.
func (e *Engine) HandleConnection(ctx context.Context, sessionCtx *common.Session) error {
observe := common.GetConnectionSetupTimeObserver(sessionCtx.Database)
if err := e.authorizeConnection(ctx); err != nil {
e.Audit.OnSessionStart(e.Context, sessionCtx, err)
return trace.Wrap(err)
}
clientConnReader := bufio.NewReader(e.clientConn)
if sessionCtx.Identity.RouteToDatabase.Username == "" {
return trace.BadParameter("database username required for Elasticsearch")
}
tlsConfig, err := e.Auth.GetTLSConfig(ctx, sessionCtx.GetExpiry(), sessionCtx.Database, sessionCtx.DatabaseUser)
if err != nil {
return trace.Wrap(err)
}
client := &http.Client{
// TODO(gavin): use an http proxy env var respecting transport
Transport: &http.Transport{
TLSClientConfig: tlsConfig,
},
}
e.Audit.OnSessionStart(e.Context, sessionCtx, nil)
defer e.Audit.OnSessionEnd(e.Context, sessionCtx)
observe()
msgFromClient := common.GetMessagesFromClientMetric(e.sessionCtx.Database)
msgFromServer := common.GetMessagesFromServerMetric(e.sessionCtx.Database)
for {
req, err := http.ReadRequest(clientConnReader)
if err != nil {
return trace.Wrap(err)
}
err = e.process(ctx, sessionCtx, req, client, msgFromClient, msgFromServer)
if err != nil {
return trace.Wrap(err)
}
}
}
// process reads request from connected elasticsearch client, processes the requests/responses and send data back
// to the client.
func (e *Engine) process(ctx context.Context, sessionCtx *common.Session, req *http.Request, client *http.Client, msgFromClient prometheus.Counter, msgFromServer prometheus.Counter) error {
msgFromClient.Inc()
if req.Body != nil {
// make sure we close the incoming request's body. ignore any close error.
defer req.Body.Close()
req.Body = io.NopCloser(utils.LimitReader(req.Body, teleport.MaxHTTPRequestSize))
}
payload, err := utils.GetAndReplaceRequestBody(req)
if err != nil {
return trace.Wrap(err)
}
copiedReq := req.Clone(ctx)
copiedReq.RequestURI = ""
copiedReq.Body = io.NopCloser(bytes.NewReader(payload))
// rewrite request URL
u, err := parseURI(sessionCtx.Database.GetURI())
if err != nil {
return trace.Wrap(err)
}
copiedReq.URL.Scheme = u.Scheme
copiedReq.URL.Host = u.Host
copiedReq.Host = u.Host
// emit an audit event regardless of failure
var responseStatusCode uint32
defer func() {
e.emitAuditEvent(copiedReq, payload, responseStatusCode, err == nil)
}()
// Send the request to elasticsearch API
resp, err := client.Do(copiedReq)
if err != nil {
return trace.Wrap(err)
}
defer resp.Body.Close()
responseStatusCode = uint32(resp.StatusCode)
msgFromServer.Inc()
return trace.Wrap(e.sendResponse(resp))
}
// emitAuditEvent writes the request and response to audit stream.
func (e *Engine) emitAuditEvent(req *http.Request, body []byte, statusCode uint32, noErr bool) {
var eventCode string
if noErr && statusCode != 0 {
eventCode = events.ElasticsearchRequestCode
} else {
eventCode = events.ElasticsearchRequestFailureCode
}
// Normally the query is passed as request body, and body content type as a header.
// Yet it can also be passed as `source` and `source_content_type` URL params, and we handle that here.
contentType := req.Header.Get("Content-Type")
source := req.URL.Query().Get("source")
if len(source) > 0 {
e.Log.InfoContext(e.Context, "'source' parameter found, overriding request body.")
body = []byte(source)
contentType = req.URL.Query().Get("source_content_type")
}
target, category := parsePath(req.URL.Path)
// Heuristic to calculate the query field.
// The priority is given to 'q' URL param. If not found, we look at the request body.
// This is not guaranteed to give us actual query, for example:
// - we may not support given API
// - we may not support given content encoding
query := req.URL.Query().Get("q")
if query == "" {
query = GetQueryFromRequestBody(e.EngineConfig, contentType, body)
}
ev := &apievents.ElasticsearchRequest{
Metadata: common.MakeEventMetadata(e.sessionCtx,
events.DatabaseSessionElasticsearchRequestEvent,
eventCode),
UserMetadata: common.MakeUserMetadata(e.sessionCtx),
SessionMetadata: common.MakeSessionMetadata(e.sessionCtx),
DatabaseMetadata: common.MakeDatabaseMetadata(e.sessionCtx),
StatusCode: statusCode,
Method: req.Method,
Path: req.URL.Path,
RawQuery: req.URL.RawQuery,
Body: body,
Headers: wrappers.Traits(req.Header),
Category: category,
Target: target,
Query: query,
}
e.Audit.EmitEvent(req.Context(), ev)
}
// sendResponse sends the response back to the elasticsearch client.
func (e *Engine) sendResponse(resp *http.Response) error {
if err := resp.Write(e.clientConn); err != nil {
return trace.Wrap(err)
}
return nil
}
// authorizeConnection does authorization check for elasticsearch connection about
// to be established.
func (e *Engine) authorizeConnection(ctx context.Context) error {
authPref, err := e.Auth.GetAuthPreference(ctx)
if err != nil {
return trace.Wrap(err)
}
state := e.sessionCtx.GetAccessState(authPref)
dbRoleMatchers := role.GetDatabaseRoleMatchers(role.RoleMatchersConfig{
Database: e.sessionCtx.Database,
DatabaseUser: e.sessionCtx.DatabaseUser,
DatabaseName: e.sessionCtx.DatabaseName,
})
err = e.sessionCtx.Checker.CheckAccess(
e.sessionCtx.Database,
state,
dbRoleMatchers...,
)
return trace.Wrap(err)
}
func parseURI(uri string) (*url.URL, error) {
if !strings.Contains(uri, "://") {
uri = "https://" + uri
}
u, err := url.Parse(uri)
if err != nil {
return nil, trace.Wrap(err)
}
// force HTTPS
u.Scheme = "https"
return u, nil
}
/*
* Teleport
* Copyright (C) 2025 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 elasticsearch
import (
"cmp"
"context"
"net"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/healthcheck"
"github.com/gravitational/teleport/lib/srv/db/healthchecks"
)
// NewHealthChecker resolves an endpoint from DB URI.
func NewHealthChecker(_ context.Context, cfg healthchecks.HealthCheckerConfig) (healthcheck.HealthChecker, error) {
dbURL, err := parseURI(cfg.Database.GetURI())
if err != nil {
return nil, trace.Wrap(err)
}
host := dbURL.Hostname()
port := cmp.Or(dbURL.Port(), "443")
hostPort := net.JoinHostPort(host, port)
return healthcheck.NewTargetDialer(func(context.Context) ([]string, error) {
return []string{hostPort}, nil
}), nil
}
/*
* 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 elasticsearch
import (
"context"
"crypto/tls"
"net"
"net/http"
"net/http/httptest"
"net/url"
"strconv"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/srv/db/common"
)
// TestServerOption allows setting test server options.
type TestServerOption func(*TestServer)
type TestServer struct {
cfg common.TestServerConfig
listener net.Listener
port string
tlsConfig *tls.Config
}
// NewTestServer returns a new instance of a test Elasticsearch server.
func NewTestServer(config common.TestServerConfig, opts ...TestServerOption) (svr *TestServer, err error) {
err = config.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
defer config.CloseOnError(&err)
tlsConfig, err := common.MakeTestServerTLSConfig(config)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig.InsecureSkipVerify = true
port, err := config.Port()
if err != nil {
return nil, trace.Wrap(err)
}
testServer := &TestServer{
cfg: config,
listener: config.Listener,
port: port,
tlsConfig: tlsConfig,
}
for _, opt := range opts {
opt(testServer)
}
return testServer, nil
}
// Serve starts serving client connections.
func (s *TestServer) Serve() error {
mux := http.NewServeMux()
mux.HandleFunc("/_sql", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Elastic-Product", "Elasticsearch")
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Content-Length", strconv.Itoa(len(testSQLResponse)))
w.Write([]byte(testSQLResponse))
})
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Elastic-Product", "Elasticsearch")
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Content-Length", strconv.Itoa(len(testQueryResponse)))
w.Write([]byte(testQueryResponse))
})
srv := &httptest.Server{
Listener: s.listener,
Config: &http.Server{Handler: mux},
TLS: s.tlsConfig,
}
srv.StartTLS()
return nil
}
func (s *TestServer) Port() string {
return s.port
}
// Close starts serving client connections.
func (s *TestServer) Close() error {
return s.listener.Close()
}
type TestClient struct {
addr string
}
func (t TestClient) Ping() (*http.Response, error) {
u := url.URL{
Scheme: "http",
Host: t.addr,
Path: "/",
}
req, err := http.NewRequest(http.MethodHead, u.String(), http.NoBody)
if err != nil {
return nil, err
}
return http.DefaultClient.Do(req)
}
func (t TestClient) Query(q string) (*http.Response, error) {
u := url.URL{
Scheme: "http",
Host: t.addr,
Path: "/_sql",
}
req, err := http.NewRequest(http.MethodGet, u.String(), strings.NewReader(q))
if err != nil {
return nil, err
}
return http.DefaultClient.Do(req)
}
// MakeTestClient returns Redis client connection according to the provided
// parameters.
func MakeTestClient(_ context.Context, config common.TestClientConfig) (*TestClient, error) {
return &TestClient{
addr: config.Address,
}, nil
}
const (
// testQueryResponse is a default successful query response.
testQueryResponse = `
{
"name" : "es01",
"cluster_name" : "docker-cluster",
"cluster_uuid" : "4--PI31pREC7AnXs_3pchA",
"version" : {
"number" : "8.3.3",
"build_flavor" : "default",
"build_type" : "docker",
"build_hash" : "801fed82df74dbe537f89b71b098ccaff88d2c56",
"build_date" : "2022-07-23T19:30:09.227964828Z",
"build_snapshot" : false,
"lucene_version" : "9.2.0",
"minimum_wire_compatibility_version" : "7.17.0",
"minimum_index_compatibility_version" : "7.0.0"
},
"tagline" : "You Know, for Search"
}`
// testSQLResponse is returned from SQL endpoint.
testSQLResponse = `{"columns":[{"name":"42","type":"integer"}],"rows":[[42]]}`
)
/*
* 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 elasticsearch
import (
"encoding/json"
"github.com/ghodss/yaml"
"github.com/gravitational/teleport/lib/srv/db/common"
)
// GetQueryFromRequestBody attempts to find the actual query from the request body, to be shown to the interested user.
func GetQueryFromRequestBody(e common.EngineConfig, contentType string, body []byte) string {
// Elasticsearch APIs have no shared schema, but the ones we support have the query either
// as 'query' or as 'knn'.
// We will attempt to deserialize the query as 'q' to discover these fields.
// The type for those is 'any': both strings and objects can be found.
var q struct {
Query any `json:"query" yaml:"query"`
Knn any `json:"knn" yaml:"knn"`
}
log := e.Log.With("content_type", contentType)
switch contentType {
// CBOR and Smile are officially supported by Elasticsearch:
// https://www.elastic.co/guide/en/elasticsearch/reference/master/api-conventions.html#_content_type_requirements
// We don't support introspection of these content types, at least for now.
case "application/cbor":
log.WarnContext(e.Context, "Content type not supported.")
return ""
case "application/smile":
log.WarnContext(e.Context, "Content type not supported.")
return ""
case "application/yaml":
if len(body) == 0 {
log.InfoContext(e.Context, "Empty request body.")
return ""
}
err := yaml.Unmarshal(body, &q)
if err != nil {
log.WarnContext(e.Context, "Error decoding request body.", "error", err)
return ""
}
case "application/json":
if len(body) == 0 {
log.InfoContext(e.Context, "Empty request body.")
return ""
}
err := json.Unmarshal(body, &q)
if err != nil {
log.WarnContext(e.Context, "Error decoding request body.", "error", err)
return ""
}
default:
log.WarnContext(e.Context, "Unknown or missing 'Content-Type', assuming 'application/json'.")
if len(body) == 0 {
log.InfoContext(e.Context, "Empty request body.")
return ""
}
err := json.Unmarshal(body, &q)
if err != nil {
log.WarnContext(e.Context, "Error decoding request body.", "error", err)
return ""
}
}
result := q.Query
if result == nil {
result = q.Knn
}
if result == nil {
return ""
}
switch qt := result.(type) {
case string:
return qt
default:
marshal, err := json.Marshal(result)
if err != nil {
log.WarnContext(e.Context, "Error encoding query to json.", "body", body, "error", err)
return ""
}
return string(marshal)
}
}
/*
* 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 protocol
import (
"bytes"
"github.com/gravitational/trace"
)
// Query represents the COM_QUERY command.
//
// https://dev.mysql.com/doc/internals/en/com-query.html
// https://mariadb.com/kb/en/com_query/
type Query struct {
packet
// query is the query text.
query string
}
// Query returns the query text.
func (p *Query) Query() string {
return p.query
}
// Quit represents the COM_QUIT command.
//
// https://dev.mysql.com/doc/internals/en/com-quit.html
// https://mariadb.com/kb/en/com_quit/
type Quit struct {
packet
}
// ChangeUser represents the COM_CHANGE_USER command.
//
// https://dev.mysql.com/doc/internals/en/com-change-user.html
// https://mariadb.com/kb/en/com_change_user/
type ChangeUser struct {
packet
// user is the requested user.
user string
}
// User returns the requested user.
func (p *ChangeUser) User() string {
return p.user
}
// schemaNamePacket is a common packet format that the packet type is followed
// by the schema name.
type schemaNamePacket struct {
packet
// schemaName is the schema name.
schemaName string
}
// SchemaName returns the schema name.
func (p *schemaNamePacket) SchemaName() string {
return p.schemaName
}
// InitDB represents the COM_INIT_DB command.
//
// COM_INIT_DB is used to specify the default schema for the connection. For
// example, "USE <schema name>" from "mysql" client sends COM_INIT_DB command
// with the schema name.
//
// https://dev.mysql.com/doc/internals/en/com-init-db.html
// https://mariadb.com/kb/en/com_init_db/
type InitDB struct {
schemaNamePacket
}
// CreateDB represents the COM_CREATE_DB command.
//
// https://dev.mysql.com/doc/internals/en/com-create-db.html
// https://mariadb.com/kb/en/com_create_db/
//
// COM_CREATE_DB creates a schema. COM_CREATE_DB is deprecated in both MySQL
// and MariaDB.
type CreateDB struct {
schemaNamePacket
}
// DropDB represents the COM_DROP_DB command.
//
// https://dev.mysql.com/doc/internals/en/com-drop-db.html
// https://mariadb.com/kb/en/com_drop_db/
//
// COM_DROP_DB drops a schema. COM_DROP_DB is deprecated in both MySQL and
// MariaDB.
type DropDB struct {
schemaNamePacket
}
// ShutDown represents the COM_SHUTDOWN command.
//
// https://dev.mysql.com/doc/internals/en/com-shutdown.html
// https://mariadb.com/kb/en/com_shutdown/
//
// COM_SHUTDOWN is used to shut down the MySQL server. COM_SHUTDOWN requires
// SHUTDOWN privileges. COM_SHUTDOWN is deprecated as of MySQL 5.7.9.
type ShutDown struct {
packet
}
// ProcessKill represents the COM_PROCESS_KILL command.
//
// https://dev.mysql.com/doc/internals/en/com-process-kill.html
// https://mariadb.com/kb/en/com_process_kill/
//
// COM_PROCESS_KILL asks the server to terminate a connection. COM_PROCESS_KILL
// is deprecated as of MySQL 5.7.11.
type ProcessKill struct {
packet
// processID is the process ID of a connection.
processID uint32
}
// ProcessID returns the process ID of a connection.
func (p *ProcessKill) ProcessID() uint32 {
return p.processID
}
// Debug represents the COM_DEBUG command.
//
// https://dev.mysql.com/doc/internals/en/com-debug.html
// https://mariadb.com/kb/en/com_debug/
//
// COM_DEBUG forces the server to dump debug information to stdout. COM_DEBUG
// requires SUPER privileges.
type Debug struct {
packet
}
// Refresh represents the COM_REFRESH command.
//
// https://dev.mysql.com/doc/internals/en/com-refresh.html
//
// COM_REFRESH calls REFRESH or FLUSH statements. COM_REFRESH is deprecated as
// of MySQL 5.7.11.
type Refresh struct {
packet
// subcommand is the string representation of the subcommand.
subcommand string
}
// Subcommand returns the string representation of the subcommand.
func (p *Refresh) Subcommand() string {
return p.subcommand
}
// parseQueryPacket parses packet bytes and returns a Packet if successful.
func parseQueryPacket(packetBytes []byte) (Packet, error) {
// Be a bit paranoid and make sure the packet is not truncated.
if len(packetBytes) < packetHeaderAndTypeSize {
return nil, trace.BadParameter("failed to parse COM_QUERY packet: %v", packetBytes)
}
// 4-byte packet header + 1-byte payload header, then query text.
return &Query{
packet: packet{bytes: packetBytes},
query: string(packetBytes[packetHeaderAndTypeSize:]),
}, nil
}
// parseQuitPacket parses packet bytes and returns a Packet if successful.
func parseQuitPacket(packetBytes []byte) (Packet, error) {
return &Quit{
packet: packet{bytes: packetBytes},
}, nil
}
// parseChangeUserPacket parses packet bytes and returns a Packet if
// successful.
func parseChangeUserPacket(packetBytes []byte) (Packet, error) {
if len(packetBytes) < packetHeaderAndTypeSize {
return nil, trace.BadParameter("failed to parse COM_CHANGE_USER packet: %v", packetBytes)
}
// User is the first null-terminated string in the payload:
// https://dev.mysql.com/doc/internals/en/com-change-user.html#packet-COM_CHANGE_USER
idx := bytes.IndexByte(packetBytes[packetHeaderAndTypeSize:], 0x00)
if idx < 0 {
return nil, trace.BadParameter("failed to parse COM_CHANGE_USER packet: %v", packetBytes)
}
return &ChangeUser{
packet: packet{bytes: packetBytes},
user: string(packetBytes[packetHeaderAndTypeSize : packetHeaderAndTypeSize+idx]),
}, nil
}
// parseSchemaNamePacket parses packet bytes and returns a schemaNamePacket if
// successful.
func parseSchemaNamePacket(packetBytes []byte) (schemaNamePacket, bool) {
unread, ok := skipHeaderAndType(packetBytes)
if !ok {
return schemaNamePacket{}, false
}
return schemaNamePacket{
packet: packet{bytes: packetBytes},
schemaName: string(unread),
}, true
}
// parseInitDBPacket parses packet bytes and returns a Packet if successful.
func parseInitDBPacket(packetBytes []byte) (Packet, error) {
parent, ok := parseSchemaNamePacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_INIT_DB packet: %v", packetBytes)
}
return &InitDB{
schemaNamePacket: parent,
}, nil
}
// parseCreateDBPacket parses packet bytes and returns a Packet if successful.
func parseCreateDBPacket(packetBytes []byte) (Packet, error) {
parent, ok := parseSchemaNamePacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_CREATE_DB packet: %v", packetBytes)
}
return &CreateDB{
schemaNamePacket: parent,
}, nil
}
// parseDropDBPacket parses packet bytes and returns a Packet if successful.
func parseDropDBPacket(packetBytes []byte) (Packet, error) {
parent, ok := parseSchemaNamePacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_DROP_DB packet: %v", packetBytes)
}
return &DropDB{
schemaNamePacket: parent,
}, nil
}
// parseShutDownPacket parses packet bytes and returns a Packet if successful.
func parseShutDownPacket(packetBytes []byte) (Packet, error) {
return &ShutDown{
packet: packet{bytes: packetBytes},
}, nil
}
// parseProcessKillPacket parses packet bytes and returns a Packet if successful.
func parseProcessKillPacket(packetBytes []byte) (Packet, error) {
unread, ok := skipHeaderAndType(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_PROCESS_KILL packet: %v", packetBytes)
}
_, processID, ok := readUint32(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_PROCESS_KILL packet: %v", packetBytes)
}
return &ProcessKill{
packet: packet{bytes: packetBytes},
processID: processID,
}, nil
}
// parseDebugPacket parses packet bytes and returns a Packet if successful.
func parseDebugPacket(packetBytes []byte) (Packet, error) {
return &Debug{
packet: packet{bytes: packetBytes},
}, nil
}
// parseRefreshPacket parses packet bytes and returns a Packet if successful.
func parseRefreshPacket(packetBytes []byte) (Packet, error) {
unread, ok := skipHeaderAndType(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_REFRESH packet: %v", packetBytes)
}
_, subcommandByte, ok := readByte(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_REFRESH packet: %v", packetBytes)
}
var subcommand string
switch subcommandByte {
case 0x01:
subcommand = "REFRESH_GRANT"
case 0x02:
subcommand = "REFRESH_LOG"
case 0x04:
subcommand = "REFRESH_TABLES"
case 0x08:
subcommand = "REFRESH_HOSTS"
case 0x10:
subcommand = "REFRESH_STATUS"
case 0x20:
subcommand = "REFRESH_THREADS"
case 0x40:
subcommand = "REFRESH_SLAVE"
case 0x80:
subcommand = "REFRESH_MASTER"
default:
return nil, trace.BadParameter("failed to parse COM_REFRESH packet: %v", packetBytes)
}
return &Refresh{
packet: packet{bytes: packetBytes},
subcommand: subcommand,
}, nil
}
/*
* 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 protocol
import (
"io"
"github.com/go-mysql-org/go-mysql/mysql"
"github.com/gravitational/trace"
)
// Packet is the common interface for MySQL wire protocol packets.
type Packet interface {
// Bytes returns the packet as raw bytes.
Bytes() []byte
}
// packet is embedded in all packets to provide common methods.
type packet struct {
// bytes is the entire packet bytes.
bytes []byte
}
// Bytes returns the packet as raw bytes.
func (p *packet) Bytes() []byte {
return p.bytes
}
// Generic represents a generic packet other than the ones recognized below.
type Generic struct {
packet
}
// ParsePacket reads a protocol packet from the connection and returns it
// in a parsed form. See ReadPacket below for the packet structure.
func ParsePacket(conn io.Reader) (Packet, error) {
packetBytes, packetType, err := ReadPacket(conn)
if err != nil {
return nil, trace.Wrap(err)
}
if int(packetType) < len(packetParsersByType) && packetParsersByType[packetType] != nil {
packet, err := packetParsersByType[packetType](packetBytes)
if err != nil {
return nil, trace.Wrap(err)
}
return packet, nil
}
return &Generic{
packet: packet{bytes: packetBytes},
}, nil
}
// ReadPacket reads a protocol packet from the connection.
//
// MySQL wire protocol packet has the following structure:
//
// 4-byte
// header payload
// ________ _________ ...
// | | |
//
// xx xx xx xx xx xx xx xx ...
//
// |_____| | |
// payload | message
// length | type
// |
// sequence
// number
//
// https://dev.mysql.com/doc/internals/en/mysql-packet.html
func ReadPacket(conn io.Reader) (pkt []byte, pktType byte, err error) {
// Read 4-byte packet header.
var header [4]byte
if _, err := io.ReadFull(conn, header[:]); err != nil {
return nil, 0, trace.ConvertSystemError(err)
}
// First 3 header bytes is the payload length, the 4th is the sequence
// number which we have no use for.
payloadLen := int(uint32(header[0]) | uint32(header[1])<<8 | uint32(header[2])<<16)
if payloadLen == 0 {
return header[:], 0, nil
}
// Read the packet payload.
// TODO(r0mant): Couple of improvements could be made here:
// * Reuse buffers instead of allocating everything from scratch every time.
// * Max payload size is 16Mb, support reading larger packets.
payload := make([]byte, payloadLen)
n, err := io.ReadFull(conn, payload)
if err != nil {
return nil, 0, trace.ConvertSystemError(err)
}
// First payload byte typically indicates the command type (query, quit,
// etc) so return it separately.
return append(header[:], payload[0:n]...), payload[0], nil
}
// WritePacket writes the provided protocol packet to the connection.
func WritePacket(pkt []byte, conn io.Writer) (int, error) {
n, err := conn.Write(pkt)
if err != nil {
return 0, trace.ConvertSystemError(err)
}
return n, nil
}
// packetParsersByType is a slice of packet parser functions by packet type.
var packetParsersByType = []func([]byte) (Packet, error){
// Server responses.
mysql.OK_HEADER: parseOKPacket,
mysql.ERR_HEADER: parseErrorPacket,
// Text protocol commands.
mysql.COM_QUERY: parseQueryPacket,
mysql.COM_QUIT: parseQuitPacket,
mysql.COM_CHANGE_USER: parseChangeUserPacket,
mysql.COM_INIT_DB: parseInitDBPacket,
mysql.COM_CREATE_DB: parseCreateDBPacket,
mysql.COM_DROP_DB: parseDropDBPacket,
mysql.COM_SHUTDOWN: parseShutDownPacket,
mysql.COM_PROCESS_KILL: parseProcessKillPacket,
mysql.COM_DEBUG: parseDebugPacket,
mysql.COM_REFRESH: parseRefreshPacket,
// Prepared statement commands.
mysql.COM_STMT_PREPARE: parseStatementPreparePacket,
mysql.COM_STMT_SEND_LONG_DATA: parseStatementSendLongDataPacket,
mysql.COM_STMT_EXECUTE: parseStatementExecutePacket,
mysql.COM_STMT_CLOSE: parseStatementClosePacket,
mysql.COM_STMT_RESET: parseStatementResetPacket,
mysql.COM_STMT_FETCH: parseStatementFetchPacket,
packetTypeStatementBulkExecute: parseStatementBulkExecutePacket,
}
const (
// packetHeaderSize is the size of the packet header.
packetHeaderSize = 4
// packetTypeSize is the size of the command type.
packetTypeSize = 1
// packetHeaderAndTypeSize is the combined size of the packet header and
// type.
packetHeaderAndTypeSize = packetHeaderSize + packetTypeSize
)
const (
// packetTypeStatementBulkExecute is a MariaDB specific packet type for
// COM_STMT_BULK_EXECUTE packets.
//
// https://mariadb.com/kb/en/com_stmt_bulk_execute/
packetTypeStatementBulkExecute = 0xfa
)
/*
* 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 protocol
import "encoding/binary"
// skipHeaderAndType skips packet header and command type, and returns rest of
// the bytes.
func skipHeaderAndType(input []byte) (unread []byte, ok bool) {
return skipBytes(input, packetHeaderAndTypeSize)
}
// skipBytes skips n bytes from input and returns rest of the bytes.
func skipBytes(input []byte, n int) (unread []byte, ok bool) {
if len(input) < n {
return nil, false
}
return input[n:], true
}
// readByte reads one byte from input and returns rest of the bytes.
func readByte(input []byte) (unread []byte, read byte, ok bool) {
if len(input) < 1 {
return nil, 0x00, false
}
return input[1:], input[0], true
}
// readUint32 reads an uint32 from input and returns rest of the bytes.
func readUint32(input []byte) (unread []byte, read uint32, ok bool) {
if len(input) < 4 {
return nil, 0, false
}
return input[4:], binary.LittleEndian.Uint32(input[:4]), true
}
// readUint16 reads an uint16 from input and returns rest of the bytes.
func readUint16(input []byte) (unread []byte, read uint16, ok bool) {
if len(input) < 2 {
return nil, 0, false
}
return input[2:], binary.LittleEndian.Uint16(input[:2]), true
}
/*
* 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 protocol
import (
"github.com/gravitational/trace"
)
// OK represents the OK packet.
//
// https://dev.mysql.com/doc/internals/en/packet-OK_Packet.html
// https://mariadb.com/kb/en/ok_packet/
type OK struct {
packet
AffectedRows uint64
}
// HasAffectedRows returns true if the packet has non-zero affected rows.
func (o *OK) HasAffectedRows() bool {
return o.packet.Bytes()[packetHeaderAndTypeSize] > 0
}
// Error represents the ERR packet.
//
// https://dev.mysql.com/doc/internals/en/packet-ERR_Packet.html
// https://mariadb.com/kb/en/err_packet/
type Error struct {
packet
// Message is the error Message
Message string
// Code is the error Code.
Code uint16
}
// Error returns the error message.
func (p *Error) Error() string {
return p.Message
}
// parseOKPacket parses packet bytes and returns a Packet if successful.
func parseOKPacket(packetBytes []byte) (Packet, error) {
return &OK{
packet: packet{bytes: packetBytes},
}, nil
}
// parseErrorPacket parses packet bytes and returns a Packet if successful.
func parseErrorPacket(packetBytes []byte) (Packet, error) {
// Figure out where in the packet the error message is.
//
// Depending on the protocol version, the packet may include additional
// fields. In protocol version 4.1 it includes '#' marker:
//
// https://dev.mysql.com/doc/internals/en/packet-ERR_Packet.html
minLen := packetHeaderSize + packetTypeSize + 2 // 4-byte header + 1-byte type + 2-byte error code
if len(packetBytes) > minLen && packetBytes[minLen] == '#' {
minLen += 6 // 1-byte marker '#' + 5-byte state
}
// Be a bit paranoid and make sure the packet is not truncated.
if len(packetBytes) < minLen {
return nil, trace.BadParameter("failed to parse ERR packet: %v", packetBytes)
}
// ignore unread bytes and "ok", we already checked len of bytes.
_, code, _ := readUint16(packetBytes[packetHeaderAndTypeSize:])
return &Error{
packet: packet{bytes: packetBytes},
Message: string(packetBytes[minLen:]),
Code: code,
}, nil
}
/*
* 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 protocol
import (
"github.com/go-mysql-org/go-mysql/mysql"
"github.com/gravitational/trace"
)
// StatementPreparePacket represents the COM_STMT_PREPARE command.
//
// https://dev.mysql.com/doc/internals/en/com-stmt-prepare.html
// https://mariadb.com/kb/en/com_stmt_prepare/
//
// COM_STMT_PREPARE creates a prepared statement from passed query string.
// Parameter placeholders are marked with "?" in the query. A COM_STMT_PREPARE
// response is expected from the server after sending this command.
type StatementPreparePacket struct {
packet
// query is the query to prepare.
query string
}
// Query returns the query text.
func (p *StatementPreparePacket) Query() string {
return p.query
}
// statementIDPacket represents a common packet format where statement ID is
// after the packet type.
//
// The statement ID is returned by the server in the COM_STMT_PREPARE response.
// All prepared statement packets except COM_STMT_PREPARE starts with the
// statement ID after the packet type to identify the prepared statement to
// use.
//
// The statement ID is an unsigned integer counter, usually starting at 1 for
// each client connection.
type statementIDPacket struct {
packet
// statementID is the ID of the associated statement.
statementID uint32
}
// StatementID returns the statement ID.
func (p *statementIDPacket) StatementID() uint32 {
return p.statementID
}
// StatementSendLongDataPacket represents the COM_STMT_SEND_LONG_DATA command.
//
// https://dev.mysql.com/doc/internals/en/com-stmt-send-long-data.html
// https://mariadb.com/kb/en/com_stmt_send_long_data/
//
// COM_STMT_SEND_LONG_DATA is used to send byte stream data to the server, and
// the server appends this data to the specified parameter upon receiving it.
// It is usually used for big blobs.
type StatementSendLongDataPacket struct {
statementIDPacket
// parameterID is the identifier of the parameter or column.
parameterID uint16
// data is the byte data sent in the packet.
data []byte
}
// ParameterID returns the parameter ID.
func (p *StatementSendLongDataPacket) ParameterID() uint16 {
return p.parameterID
}
// Data returns the data in bytes.
func (p *StatementSendLongDataPacket) Data() []byte {
return p.data
}
// StatementExecutePacket represents the COM_STMT_EXECUTE command.
//
// https://dev.mysql.com/doc/internals/en/com-stmt-execute.html
// https://mariadb.com/kb/en/com_stmt_execute/
//
// COM_STMT_EXECUTE asks the server to execute a prepared statement, with the
// types and values for the placeholders.
//
// Statement ID "-1" (0xffffffff) can be used to indicate the last statement
// prepared on current connection, for MariaDB server version 10.2 and above.
type StatementExecutePacket struct {
statementIDPacket
// cursorFlag specifies type of the cursor.
cursorFlag byte
// iterations is the iteration count specified in the command. The MySQL
// doc states that it is always 1.
iterations uint32
// nullBitmapAndParameters are raw packet bytes that represent a null
// bitmap and parameters with types and values. They are not decoded in the
// initial parsing because number of parameters is unknown.
nullBitmapAndParameters []byte
}
// Parameters returns a slice of parameters.
func (p *StatementExecutePacket) Parameters(definitions []mysql.Field) (parameters []any, ok bool) {
// TODO(greedy52) implement parsing of null bitmap, parameter types, and
// paramerter binary values.
return nil, true
}
// StatementClosePacket represents the COM_STMT_CLOSE command.
//
// https://dev.mysql.com/doc/internals/en/com-stmt-close.html
// https://mariadb.com/kb/en/3-binary-protocol-prepared-statements-com_stmt_close/
//
// COM_STMT_CLOSE deallocates a prepared statement.
type StatementClosePacket struct {
statementIDPacket
}
// StatementResetPacket represents the COM_STMT_RESET command.
//
// https://dev.mysql.com/doc/internals/en/com-stmt-reset.html
// https://mariadb.com/kb/en/com_stmt_reset/
//
// COM_STMT_RESET resets the data of a prepared statement which was accumulated
// with COM_STMT_SEND_LONG_DATA.
type StatementResetPacket struct {
statementIDPacket
}
// StatementFetchPacket represents the COM_STMT_FETCH command.
//
// https://dev.mysql.com/doc/internals/en/com-stmt-fetch.html
// https://mariadb.com/kb/en/com_stmt_fetch/
//
// COM_STMT_FETCH fetch rows from a existing resultset after a
// COM_STMT_EXECUTE.
type StatementFetchPacket struct {
statementIDPacket
// rowsCount number of rows to fetch.
rowsCount uint32
}
// RowsCount returns number of rows to fetch.
func (s *StatementFetchPacket) RowsCount() uint32 {
return s.rowsCount
}
// StatementBulkExecutePacket represents the COM_STMT_BULK_EXECUTE command.
//
// https://mariadb.com/kb/en/com_stmt_bulk_execute/
//
// COM_STMT_BULK_EXECUTE executes a bulk insert of a previously prepared
// statement.
type StatementBulkExecutePacket struct {
statementIDPacket
// bulkFlag is a flag specifies either 64 (return generated auto-increment
// IDs) or 128 (send types to server).
bulkFlag uint16
// parameters are raw packet bytes that contain parameter type and values.
// They are not decoded in the initial parsing because number of parameters
// is unknown.
parameters []byte
}
// Parameters returns a slice of parameters.
func (p *StatementBulkExecutePacket) Parameters(definitions []mysql.Field) (parameters []any, ok bool) {
// TODO(greedy52) implement parsing of parameters from
// COM_STMT_BULK_EXECUTE packet.
return nil, true
}
// parseStatementPreparePacket parses packet bytes and returns a Packet if
// successful.
func parseStatementPreparePacket(packetBytes []byte) (Packet, error) {
unread, ok := skipHeaderAndType(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_PREPARE packet: %v", packetBytes)
}
return &StatementPreparePacket{
packet: packet{bytes: packetBytes},
query: string(unread),
}, nil
}
// parseStatementIDPacket parses packet bytes and returns a statementIDPacket
// if successful.
func parseStatementIDPacket(packetBytes []byte) (statementIDPacket, []byte, bool) {
unread, ok := skipHeaderAndType(packetBytes)
if !ok {
return statementIDPacket{}, nil, false
}
unread, statementID, ok := readUint32(unread)
if !ok {
return statementIDPacket{}, nil, false
}
return statementIDPacket{
packet: packet{bytes: packetBytes},
statementID: statementID,
}, unread, true
}
// parseStatementSendLongDataPacket parses packet bytes and returns a Packet if
// successful.
func parseStatementSendLongDataPacket(packetBytes []byte) (Packet, error) {
parent, unread, ok := parseStatementIDPacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_SEND_LONG_DATA packet: %v", packetBytes)
}
unread, parameterID, ok := readUint16(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_SEND_LONG_DATA packet: %v", packetBytes)
}
return &StatementSendLongDataPacket{
statementIDPacket: parent,
parameterID: parameterID,
data: unread,
}, nil
}
// parseStatementExecutePacket parses packet bytes and returns a Packet if
// successful.
func parseStatementExecutePacket(packetBytes []byte) (Packet, error) {
parent, unread, ok := parseStatementIDPacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_EXECUTE packet: %v", packetBytes)
}
unread, cursorFlag, ok := readByte(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_EXECUTE packet: %v", packetBytes)
}
unread, iterations, ok := readUint32(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_EXECUTE packet: %v", packetBytes)
}
return &StatementExecutePacket{
statementIDPacket: parent,
cursorFlag: cursorFlag,
iterations: iterations,
nullBitmapAndParameters: unread,
}, nil
}
// parseStatementClosePacket parses packet bytes and returns a Packet if
// successful.
func parseStatementClosePacket(packetBytes []byte) (Packet, error) {
parent, _, ok := parseStatementIDPacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_CLOSE packet: %v", packetBytes)
}
return &StatementClosePacket{
statementIDPacket: parent,
}, nil
}
// parseStatementResetPacket parses packet bytes and returns a Packet if
// successful.
func parseStatementResetPacket(packetBytes []byte) (Packet, error) {
parent, _, ok := parseStatementIDPacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_RESET packet: %v", packetBytes)
}
return &StatementResetPacket{
statementIDPacket: parent,
}, nil
}
// parseStatementFetchPacket parses packet bytes and returns a Packet if
// successful.
func parseStatementFetchPacket(packetBytes []byte) (Packet, error) {
parent, unread, ok := parseStatementIDPacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_FETCH packet: %v", packetBytes)
}
_, rowsCount, ok := readUint32(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_FETCH packet: %v", packetBytes)
}
return &StatementFetchPacket{
statementIDPacket: parent,
rowsCount: rowsCount,
}, nil
}
// parseStatementBulkExecutePacket parses packet bytes and returns a Packet if
// successful.
func parseStatementBulkExecutePacket(packetBytes []byte) (Packet, error) {
parent, unread, ok := parseStatementIDPacket(packetBytes)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_BULK_EXECUTE packet: %v", packetBytes)
}
unread, bulkFlag, ok := readUint16(unread)
if !ok {
return nil, trace.BadParameter("failed to parse COM_STMT_BULK_EXECUTE packet: %v", packetBytes)
}
return &StatementBulkExecutePacket{
statementIDPacket: parent,
bulkFlag: bulkFlag,
parameters: unread,
}, nil
}
/*
* 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 protocol
import (
"bufio"
"bytes"
"context"
"errors"
"io"
"net"
"time"
"github.com/go-mysql-org/go-mysql/client"
"github.com/go-mysql-org/go-mysql/mysql"
mysqlpacket "github.com/go-mysql-org/go-mysql/packet"
"github.com/gravitational/trace"
)
// FetchMySQLVersionInternal connects to a MySQL instance with provided dialer and tries to read the server
// version from initial handshake message. Error is returned in case of connection failure or when MySQL
// returns ERR package.
func FetchMySQLVersionInternal(ctx context.Context, dialer client.Dialer, databaseURI string) (string, error) {
conn, err := dialer(ctx, "tcp", databaseURI)
if err != nil {
return "", trace.ConnectionProblem(err, "failed to connect to MySQL")
}
defer conn.Close()
return ReadMySQLVersion(ctx, conn)
}
// ReadMySQLVersion tries to read the server version from initial handshake message.
// Error is returned if MySQL returns ERR package.
func ReadMySQLVersion(ctx context.Context, conn net.Conn) (string, error) {
// Set connection deadline if passed context has it.
if deadline, ok := ctx.Deadline(); ok {
if err := conn.SetReadDeadline(deadline); err != nil {
return "", trace.Wrap(err)
}
defer conn.SetReadDeadline(time.Time{})
}
connBuf := NewBufferedConn(ctx, conn)
pkgType, err := connBuf.Peek(5)
if err != nil {
return "", trace.Wrap(err)
}
// ref: https://dev.mysql.com/doc/internals/en/mysql-packet.html
// https://dev.mysql.com/doc/internals/en/packet-ERR_Packet.html
if pkgType[4] == mysql.ERR_HEADER {
return readHandshakeError(connBuf)
}
return readHandshakeServerVersion(connBuf)
}
// readHandshakeServerVersion reads MySQL initial handshake message and returns the server version.
func readHandshakeServerVersion(connBuf net.Conn) (string, error) {
dbConn := mysqlpacket.NewTLSConn(connBuf)
handshake, err := dbConn.ReadPacket()
if err != nil {
return "", trace.ConnectionProblem(err, "failed to read the MySQL handshake")
}
if len(handshake) == 0 {
return "", trace.Errorf("server returned empty handshake packet")
}
// ref: https://dev.mysql.com/doc/internals/en/connection-phase-packets.html#packet-Protocol::Handshake
versionLength := bytes.IndexByte(handshake[1:], 0x00)
if versionLength == -1 {
return "", trace.Errorf("failed to read the MySQL server version")
}
return string(handshake[1 : 1+versionLength]), nil
}
// readHandshakeError reads and returns an error message from
func readHandshakeError(connBuf io.Reader) (string, error) {
handshakePacket, err := ParsePacket(connBuf)
if err != nil {
return "", err
}
errPackage, ok := handshakePacket.(*Error)
if !ok {
return "", trace.BadParameter("expected MySQL error package, got %T", handshakePacket)
}
return "", trace.ConnectionProblem(errors.New("failed to fetch MySQL version"), "%s", errPackage.Error())
}
// IsHandshakeV10Packet peeks into the conn and checks for a handshake v10 packet.
// The results of this function are only meaningful during the connection phase
// of the MySQL protocol. It is the caller's responsibility to only use this
// function during the connection phase.
// https://dev.mysql.com/doc/dev/mysql-server/latest/page_protocol_connection_phase_packets_protocol_handshake_v10.html
func IsHandshakeV10Packet(conn BufferedConn) (bool, error) {
pkgHeaderAndType, err := conn.Peek(packetHeaderAndTypeSize)
if err != nil {
return false, trace.Wrap(err)
}
const typeIdx = packetHeaderAndTypeSize - 1
return pkgHeaderAndType[typeIdx] == 10, nil
}
// BufferedConn is a net.Conn wrapper with additional Peek() method.
type BufferedConn struct {
net.Conn
ctx context.Context
reader *bufio.Reader
}
// NewBufferedConn wraps a [net.Conn] in a new [BufferedConn].
func NewBufferedConn(ctx context.Context, conn net.Conn) BufferedConn {
return BufferedConn{
ctx: ctx,
reader: bufio.NewReader(conn),
Conn: conn,
}
}
// Discard discards n bytes from the reader.
// It's basically a wrapper around (bufio.Reader).Discard()
func (b BufferedConn) Discard(n int) (discarded int, err error) {
if err := b.ctx.Err(); err != nil {
return 0, err
}
return b.reader.Discard(n)
}
// Peek reads n bytes without advancing the reader.
// It's basically a wrapper around (bufio.Reader).Peek()
func (b BufferedConn) Peek(n int) ([]byte, error) {
if err := b.ctx.Err(); err != nil {
return nil, err
}
return b.reader.Peek(n)
}
// Read returns data from underlying buffer.
func (b BufferedConn) Read(p []byte) (int, error) {
if err := b.ctx.Err(); err != nil {
return 0, err
}
return b.reader.Read(p)
}
/*
* 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 opensearch
import (
"strings"
apievents "github.com/gravitational/teleport/api/types/events"
)
// parsePath returns (optional) target of query as well as the event category.
func parsePath(path string) (string, apievents.OpenSearchCategory) {
parts := strings.Split(path, "/")
// empty string or lone /
if len(parts) < 2 {
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_GENERAL
}
// underscore
if strings.HasPrefix(parts[1], "_") {
switch parts[1] {
case
// search
"_search", // https://opensearch.org/docs/latest/api-reference/search/
"_count", // https://opensearch.org/docs/latest/api-reference/count/
"_msearch", // https://opensearch.org/docs/2.6/api-reference/multi-search/
"_render", // https://opensearch.org/docs/2.6/search-plugins/search-template/
"_validate",
"_field_caps",
"_rank_eval",
"_search_shards":
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_SEARCH
case "_plugins":
// handle plugins
// length check
if len(parts) < 3 {
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_GENERAL
}
switch parts[2] {
case
"_sql", // https://opensearch.org/docs/2.6/search-plugins/sql/sql-ppl-api/
"_ppl": // https://opensearch.org/docs/2.6/search-plugins/sql/sql-ppl-api/
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_SQL
case
"_asynchronous_search", // https://opensearch.org/docs/2.6/search-plugins/async/index/
"_knn": // https://opensearch.org/docs/2.6/search-plugins/knn/api/
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_SEARCH
case "_security": // https://opensearch.org/docs/2.6/security/index/
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_SECURITY
}
case "_all":
// fall through
default:
// starts with _, but we don't have logic to handle it in a special way.
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_GENERAL
}
}
// length check
if len(parts) < 3 {
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_GENERAL
}
// a number of APIs are invoked by providing a target first, e.g. /<target>/_search, where <target> is an index or expression matching a group of indices.
switch parts[2] {
case
// search variants
"_search", // https://opensearch.org/docs/2.6/api-reference/search/
"_count", // https://opensearch.org/docs/2.6/api-reference/count/
"_msearch", // https://opensearch.org/docs/2.6/api-reference/multi-search/
"_explain", // https://opensearch.org/docs/2.6/api-reference/explain/
"_rank_eval", // https://opensearch.org/docs/2.6/api-reference/rank-eval/
"_search_shards",
"_validate",
"_field_caps",
"_terms_enum":
return parts[1], apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_SEARCH
}
// no special handling, general case.
return "", apievents.OpenSearchCategory_OPEN_SEARCH_CATEGORY_GENERAL
}
/*
* 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 opensearch
import (
"bufio"
"bytes"
"context"
"encoding/json"
"io"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"github.com/gravitational/trace"
"github.com/prometheus/client_golang/prometheus"
"github.com/gravitational/teleport"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/api/types/wrappers"
"github.com/gravitational/teleport/lib/cloud/awsconfig"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/srv/db/common"
"github.com/gravitational/teleport/lib/srv/db/common/role"
"github.com/gravitational/teleport/lib/srv/db/elasticsearch"
"github.com/gravitational/teleport/lib/utils"
libaws "github.com/gravitational/teleport/lib/utils/aws"
)
// NewEngine create new OpenSearch engine.
func NewEngine(ec common.EngineConfig) common.Engine {
return &Engine{
EngineConfig: ec,
}
}
// Engine handles connections from OpenSearch clients coming from Teleport
// proxy over reverse tunnel.
type Engine struct {
// EngineConfig is the common database engine configuration.
common.EngineConfig
// clientConn is a client connection.
clientConn net.Conn
// sessionCtx is current session context.
sessionCtx *common.Session
}
// InitializeConnection initializes the engine with the client connection.
func (e *Engine) InitializeConnection(clientConn net.Conn, sessionCtx *common.Session) error {
e.clientConn = clientConn
e.sessionCtx = sessionCtx
return nil
}
// errorDetails contains error details.
type errorDetails struct {
Reason string `json:"reason"`
Type string `json:"type"`
}
// errorResponse will be returned to the client in case of error.
type errorResponse struct {
Error errorDetails `json:"error"`
Status int `json:"status"`
}
// SendError sends an error to OpenSearch client.
func (e *Engine) SendError(err error) {
if e.clientConn == nil || err == nil || utils.IsOKNetworkError(err) {
return
}
cause := errorResponse{
Error: errorDetails{
Reason: err.Error(),
Type: "internal_server_error_exception",
},
Status: http.StatusInternalServerError,
}
// Different error for access denied case.
if trace.IsAccessDenied(err) {
cause.Status = http.StatusUnauthorized
cause.Error.Type = "access_denied_exception"
}
jsonBody, err := json.Marshal(cause)
if err != nil {
e.Log.ErrorContext(e.Context, "Failed to marshal error response.", "error", err)
return
}
response := &http.Response{
ProtoMajor: 1,
ProtoMinor: 1,
StatusCode: cause.Status,
Body: io.NopCloser(bytes.NewBuffer(jsonBody)),
Header: map[string][]string{
"Content-Type": {"application/json"},
"Content-Length": {strconv.Itoa(len(jsonBody))},
},
}
if err := response.Write(e.clientConn); err != nil {
e.Log.ErrorContext(e.Context, "Failed to send an error to the client.", "error", err)
return
}
}
// HandleConnection authorizes the incoming client connection, connects to the
// target OpenSearch server and starts proxying requests between client/server.
func (e *Engine) HandleConnection(ctx context.Context, _ *common.Session) error {
observe := common.GetConnectionSetupTimeObserver(e.sessionCtx.Database)
err := e.checkAccess(ctx)
if err != nil {
e.Audit.OnSessionStart(e.Context, e.sessionCtx, err)
return trace.Wrap(err)
}
// TODO(Tener):
// Consider rewriting to support HTTP2 clients.
// Ideally we should have shared middleware for DB clients using HTTP, instead of separate per-engine implementations.
tr, err := e.getTransport(ctx)
if err != nil {
return trace.Wrap(err)
}
e.Audit.OnSessionStart(e.Context, e.sessionCtx, nil)
defer e.Audit.OnSessionEnd(e.Context, e.sessionCtx)
clientConnReader := bufio.NewReader(e.clientConn)
observe()
msgFromClient := common.GetMessagesFromClientMetric(e.sessionCtx.Database)
msgFromServer := common.GetMessagesFromServerMetric(e.sessionCtx.Database)
for {
req, err := http.ReadRequest(clientConnReader)
if err != nil {
return trace.Wrap(err)
}
if err := e.process(ctx, tr, req, msgFromClient, msgFromServer); err != nil {
return trace.Wrap(err)
}
}
}
// process reads request from connected OpenSearch client, processes the requests/responses and send data back
// to the client.
func (e *Engine) process(ctx context.Context, tr *http.Transport, req *http.Request, msgFromClient prometheus.Counter, msgFromServer prometheus.Counter) error {
msgFromClient.Inc()
if req.Body != nil {
// make sure we close the incoming request's body. ignore any close error.
defer req.Body.Close()
req.Body = io.NopCloser(utils.LimitReader(req.Body, teleport.MaxHTTPRequestSize))
}
reqCopy, payload, err := e.rewriteRequest(ctx, req)
if err != nil {
return trace.Wrap(err)
}
// emit an audit event regardless of failure
var responseStatusCode uint32
defer func() {
e.emitAuditEvent(reqCopy, payload, responseStatusCode, err == nil)
}()
signedReq, err := e.getSignedRequest(reqCopy)
if err != nil {
return trace.Wrap(err)
}
//nolint:bodyclose // resp will be closed in sendResponse().
resp, err := tr.RoundTrip(signedReq)
if err != nil {
return trace.Wrap(err)
}
responseStatusCode = uint32(resp.StatusCode)
msgFromServer.Inc()
return trace.Wrap(e.sendResponse(resp))
}
func (e *Engine) getTransport(ctx context.Context) (*http.Transport, error) {
tr, err := defaults.Transport()
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig, err := e.Auth.GetTLSConfig(ctx, e.sessionCtx.GetExpiry(), e.sessionCtx.Database, e.sessionCtx.DatabaseUser)
if err != nil {
return nil, trace.Wrap(err)
}
tr.TLSClientConfig = tlsConfig
return tr, nil
}
func (e *Engine) getSignedRequest(reqCopy *http.Request) (*http.Request, error) {
roleArn, err := libaws.BuildRoleARN(e.sessionCtx.DatabaseUser, e.sessionCtx.Database.GetAWS().Region, e.sessionCtx.Database.GetAWS().AccountID)
if err != nil {
return nil, trace.Wrap(err)
}
meta := e.sessionCtx.Database.GetAWS()
awsCfg, err := e.AWSConfigProvider.GetConfig(e.Context, meta.Region,
awsconfig.WithAssumeRole(meta.AssumeRoleARN, meta.ExternalID),
awsconfig.WithDetailedAssumeRole(awsconfig.AssumeRole{
RoleARN: roleArn,
ExternalID: meta.ExternalID,
SessionName: e.sessionCtx.Identity.Username,
}),
awsconfig.WithAmbientCredentials(),
)
if err != nil {
return nil, trace.Wrap(err)
}
signCtx := &libaws.SigningCtx{
Clock: e.Clock,
Credentials: awsCfg.Credentials,
SigningName: "es",
SigningRegion: e.sessionCtx.Database.GetAWS().Region,
}
signedReq, err := libaws.SignRequest(e.Context, reqCopy, signCtx)
if err != nil {
return nil, trace.Wrap(err)
}
return signedReq, nil
}
func (e *Engine) rewriteRequest(ctx context.Context, req *http.Request) (*http.Request, []byte, error) {
payload, err := utils.GetAndReplaceRequestBody(req)
if err != nil {
return nil, nil, trace.Wrap(err)
}
reqCopy := req.Clone(ctx)
reqCopy.RequestURI = ""
reqCopy.Body = io.NopCloser(bytes.NewReader(payload))
// Connection is hop-by-hop header, drop.
reqCopy.Header.Del("Connection")
// rewrite request URL
u, err := parseURI(e.sessionCtx.Database.GetURI())
if err != nil {
return nil, nil, trace.Wrap(err)
}
reqCopy.URL.Scheme = u.Scheme
reqCopy.URL.Host = u.Host
reqCopy.Host = u.Host
return reqCopy, payload, nil
}
// emitAuditEvent writes the request and response to audit stream.
func (e *Engine) emitAuditEvent(req *http.Request, body []byte, statusCode uint32, noErr bool) {
var eventCode string
if noErr && statusCode != 0 {
eventCode = events.OpenSearchRequestCode
} else {
eventCode = events.OpenSearchRequestFailureCode
}
// Normally the query is passed as request body, and body content type as a header.
// Yet it can also be passed as `source` and `source_content_type` URL params, and we handle that here.
contentType := req.Header.Get("Content-Type")
source := req.URL.Query().Get("source")
if len(source) > 0 {
e.Log.InfoContext(e.Context, "'source' parameter found, overriding request body.")
body = []byte(source)
contentType = req.URL.Query().Get("source_content_type")
}
target, category := parsePath(req.URL.Path)
// Heuristic to calculate the query field.
// The priority is given to 'q' URL param. If not found, we look at the request body.
// This is not guaranteed to give us actual query, for example:
// - we may not support given API
// - we may not support given content encoding
query := req.URL.Query().Get("q")
if query == "" {
query = elasticsearch.GetQueryFromRequestBody(e.EngineConfig, contentType, body)
}
ev := &apievents.OpenSearchRequest{
Metadata: common.MakeEventMetadata(e.sessionCtx,
events.DatabaseSessionOpenSearchRequestEvent,
eventCode),
UserMetadata: common.MakeUserMetadata(e.sessionCtx),
SessionMetadata: common.MakeSessionMetadata(e.sessionCtx),
DatabaseMetadata: common.MakeDatabaseMetadata(e.sessionCtx),
StatusCode: statusCode,
Method: req.Method,
Path: req.URL.Path,
RawQuery: req.URL.RawQuery,
Body: body,
Headers: wrappers.Traits(req.Header),
Category: category,
Target: target,
Query: query,
}
e.Audit.EmitEvent(e.Context, ev)
}
// sendResponse sends the response back to the OpenSearch client.
func (e *Engine) sendResponse(serverResponse *http.Response) error {
if serverResponse.Body != nil {
defer serverResponse.Body.Close()
serverResponse.Body = io.NopCloser(io.LimitReader(serverResponse.Body, teleport.MaxHTTPResponseSize))
}
payload, err := utils.GetAndReplaceResponseBody(serverResponse)
if err != nil {
return trace.Wrap(err)
}
// serverResponse may be HTTP2 response, but we should reply with HTTP 1.1
clientResponse := &http.Response{
ProtoMajor: 1,
ProtoMinor: 1,
StatusCode: serverResponse.StatusCode,
Body: io.NopCloser(bytes.NewBuffer(payload)),
Header: serverResponse.Header.Clone(),
ContentLength: int64(len(payload)),
}
if err := clientResponse.Write(e.clientConn); err != nil {
return trace.Wrap(err)
}
return nil
}
// checkAccess does authorization check for OpenSearch connection about
// to be established.
func (e *Engine) checkAccess(ctx context.Context) error {
if e.sessionCtx.Identity.RouteToDatabase.Username == "" {
return trace.BadParameter("database username required for OpenSearch")
}
authPref, err := e.Auth.GetAuthPreference(ctx)
if err != nil {
return trace.Wrap(err)
}
state := e.sessionCtx.GetAccessState(authPref)
dbRoleMatchers := role.GetDatabaseRoleMatchers(role.RoleMatchersConfig{
Database: e.sessionCtx.Database,
DatabaseUser: e.sessionCtx.DatabaseUser,
DatabaseName: e.sessionCtx.DatabaseName,
})
err = e.sessionCtx.Checker.CheckAccess(
e.sessionCtx.Database,
state,
dbRoleMatchers...,
)
return trace.Wrap(err)
}
func parseURI(uri string) (*url.URL, error) {
if !strings.Contains(uri, "://") {
uri = "https://" + uri
}
u, err := url.Parse(uri)
if err != nil {
return nil, trace.Wrap(err)
}
// force HTTPS
u.Scheme = "https"
return u, nil
}
/*
* Teleport
* Copyright (C) 2025 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 opensearch
import (
"cmp"
"context"
"net"
"github.com/gravitational/trace"
"golang.org/x/net/http/httpproxy"
"github.com/gravitational/teleport/lib/healthcheck"
"github.com/gravitational/teleport/lib/srv/db/healthchecks"
)
// NewHealthChecker resolves an endpoint from DB URI.
func NewHealthChecker(_ context.Context, cfg healthchecks.HealthCheckerConfig) (healthcheck.HealthChecker, error) {
dbURL, err := parseURI(cfg.Database.GetURI())
if err != nil {
return nil, trace.Wrap(err)
}
// Not all of our DB engines respect http proxy env vars, but this one does.
// The endpoint resolved for TCP health checks should be the one that the
// agent will actually connect to, since often proxy env vars are set to
// accommodate self-imposed network restrictions that force external traffic
// to go through a proxy.
proxyFunc := httpproxy.FromEnvironment().ProxyFunc()
proxyURL, err := proxyFunc(dbURL)
if err != nil {
return nil, trace.Wrap(err)
}
if proxyURL != nil {
dbURL = proxyURL
}
host := dbURL.Hostname()
port := cmp.Or(dbURL.Port(), "443")
hostPort := net.JoinHostPort(host, port)
return healthcheck.NewTargetDialer(func(context.Context) ([]string, error) {
return []string{hostPort}, nil
}), nil
}
/*
* 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 opensearch
import (
"context"
"crypto/tls"
"log/slog"
"net"
"net/http"
"net/http/httptest"
"strconv"
"github.com/gravitational/trace"
"github.com/opensearch-project/opensearch-go/v2"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/srv/db/common"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/log/logtest"
)
// TestServerOption allows setting test server options.
type TestServerOption func(*TestServer)
type TestServer struct {
cfg common.TestServerConfig
listener net.Listener
port string
tlsConfig *tls.Config
log *slog.Logger
}
// NewTestServer returns a new instance of a test OpenSearch server.
func NewTestServer(config common.TestServerConfig, opts ...TestServerOption) (svr *TestServer, err error) {
err = config.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
defer config.CloseOnError(&err)
tlsConfig, err := common.MakeTestServerTLSConfig(config)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig.InsecureSkipVerify = true
port, err := config.Port()
if err != nil {
return nil, trace.Wrap(err)
}
testServer := &TestServer{
cfg: config,
listener: config.Listener,
port: port,
tlsConfig: tlsConfig,
log: logtest.With(
teleport.ComponentKey, defaults.ProtocolOpenSearch,
"name", config.Name,
),
}
for _, opt := range opts {
opt(testServer)
}
return testServer, nil
}
// Serve starts serving client connections.
func (s *TestServer) Serve() error {
mux := http.NewServeMux()
mux.HandleFunc("/_count", func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Content-Length", strconv.Itoa(len(testCountResponse)))
w.Write([]byte(testCountResponse))
})
mux.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) {
s.log.DebugContext(r.Context(), "Handling request", "url", logutils.StringerAttr(r.URL))
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Content-Length", strconv.Itoa(len(testQueryResponse)))
w.Write([]byte(testQueryResponse))
})
srv := &httptest.Server{
Listener: s.listener,
Config: &http.Server{Handler: mux},
TLS: s.tlsConfig,
}
srv.StartTLS()
return nil
}
func (s *TestServer) Port() string {
return s.port
}
// Close starts serving client connections.
func (s *TestServer) Close() error {
return s.listener.Close()
}
// MakeTestClient returns Redis client connection according to the provided
// parameters.
func MakeTestClient(_ context.Context, config common.TestClientConfig) (*opensearch.Client, error) {
clt, err := opensearch.NewClient(opensearch.Config{
Addresses: []string{"http://" + config.Address},
})
return clt, trace.Wrap(err)
}
const (
// testQueryResponse is a default successful query response.
testQueryResponse = `
{
"name" : "5b1234425a372fe585532acf1da64323",
"cluster_name" : "123476220453:test",
"cluster_uuid" : "G_uZbXf-TgKJaGU8sEc70g",
"version" : {
"distribution" : "opensearch",
"number" : "2.3.0",
"build_type" : "tar",
"build_hash" : "unknown",
"build_date" : "2022-11-10T22:04:34.357368Z",
"build_snapshot" : false,
"lucene_version" : "9.3.0",
"minimum_wire_compatibility_version" : "7.10.0",
"minimum_index_compatibility_version" : "7.0.0"
},
"tagline" : "The OpenSearch Project: https://opensearch.org/"
}`
// testCountResponse is returned from Count endpoint.
testCountResponse = `{"count":31874,"_shards":{"total":6,"successful":6,"skipped":0,"failed":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 protocol
import (
"bytes"
"encoding/binary"
"io"
"github.com/gravitational/trace"
mssql "github.com/microsoft/go-mssqldb"
)
// Login7Packet represents a Login7 packet that defines authentication rules
// between the client and the server.
//
// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/773a62b6-ee89-4c02-9e5e-344882630aac
type Login7Packet struct {
packet Packet
header Login7Header
username string
database string
}
// Username returns the username from the Login7 packet.
func (p *Login7Packet) Username() string {
return p.username
}
// Database returns the database from the Login7 packet. May be empty.
func (p *Login7Packet) Database() string {
return p.database
}
// OptionFlags1 returns the packet's first set of option flags.
func (p *Login7Packet) OptionFlags1() uint8 {
return p.header.OptionFlags1
}
// OptionFlags2 returns the packet's second set of option flags.
func (p *Login7Packet) OptionFlags2() uint8 {
return p.header.OptionFlags2
}
// TypeFlags returns the packet's set of type flags.
func (p *Login7Packet) TypeFlags() uint8 {
return p.header.TypeFlags
}
// PacketSize return the packet size from the Login7 packet.
// Packet size is used by a server to negation the size of max packet length.
func (p *Login7Packet) PacketSize() uint16 {
return uint16(p.header.PacketSize)
}
// Login7Header contains options and offset/length pairs parsed from the Login7
// packet sent by client.
//
// Note: the order of fields in the struct matters as it gets unpacked from the
// binary stream.
type Login7Header struct {
Length uint32
TDSVersion uint32
PacketSize uint32
ClientProgVer uint32
ClientPID uint32
ConnectionID uint32
OptionFlags1 uint8
OptionFlags2 uint8
TypeFlags uint8
OptionFlags3 uint8
ClientTimezone int32
ClientLCID uint32
IbHostName uint16 // offset
CchHostName uint16 // length
IbUserName uint16
CchUserName uint16
IbPassword uint16
CchPassword uint16
IbAppName uint16
CchAppName uint16
IbServerName uint16
CchServerName uint16
IbUnused uint16
CbUnused uint16
IbCltIntName uint16
CchCltIntName uint16
IbLanguage uint16
CchLanguage uint16
IbDatabase uint16
CchDatabase uint16
ClientID [6]byte
IbSSPI uint16
CbSSPI uint16
IbAtchDBFile uint16
CchAtchDBFile uint16
IbChangePassword uint16
CchChangePassword uint16
CbSSPILong uint32
}
// ReadLogin7Packet reads Login7 packet from the reader.
func ReadLogin7Packet(r io.Reader) (*Login7Packet, error) {
pkt, err := ReadPacket(r)
if err != nil {
return nil, trace.Wrap(err)
}
if pkt.Type() != PacketTypeLogin7 {
return nil, trace.BadParameter("expected Login7 packet, got: %#v", pkt)
}
var header Login7Header
if err := binary.Read(bytes.NewReader(pkt.Data()), binary.LittleEndian, &header); err != nil {
return nil, trace.Wrap(err)
}
username, err := readUsername(pkt, header)
if err != nil {
return nil, trace.Wrap(err)
}
database, err := readDatabase(pkt, header)
if err != nil {
return nil, trace.Wrap(err)
}
return &Login7Packet{
packet: pkt,
header: header,
username: username,
database: database,
}, nil
}
// errInvalidPacket is returned when Login7 package contains invalid data.
var errInvalidPacket = trace.Errorf("invalid login7 packet")
// readUsername reads username from login7 package.
func readUsername(pkt Packet, header Login7Header) (string, error) {
if len(pkt.Data()) < int(header.IbUserName)+int(header.CchUserName)*2 {
return "", errInvalidPacket
}
// Decode username and database from the packet. Offset/length are counted
// from the beginning of entire packet data (excluding header).
username, err := mssql.ParseUCS2String(
pkt.Data()[header.IbUserName : header.IbUserName+header.CchUserName*2])
if err != nil {
return "", trace.Wrap(err)
}
return username, nil
}
// readDatabase reads database name from login7 package.
func readDatabase(pkt Packet, header Login7Header) (string, error) {
if len(pkt.Data()) < int(header.IbDatabase)+int(header.CchDatabase)*2 {
return "", errInvalidPacket
}
database, err := mssql.ParseUCS2String(
pkt.Data()[header.IbDatabase : header.IbDatabase+header.CchDatabase*2])
if err != nil {
return "", trace.Wrap(err)
}
return database, nil
}
/*
* 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 protocol
import (
"bytes"
"encoding/binary"
"io"
"runtime/debug"
"github.com/gravitational/trace"
)
// PacketHeader represents a 8-byte packet header.
//
// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/7af53667-1b72-4703-8258-7984e838f746
//
// Note: the order of fields in the struct matters as it gets unpacked from the
// binary stream.
type PacketHeader struct {
Type uint8
Status uint8
Length uint16 // network byte order (big-endian)
SPID uint16 // network byte order (big-endian)
PacketID uint8
Window uint8
}
// Marshal marshals the packet header to the wire protocol byte representation.
func (h *PacketHeader) Marshal() ([]byte, error) {
buf := bytes.NewBuffer(make([]byte, 0, packetHeaderSize))
if err := binary.Write(buf, binary.BigEndian, h); err != nil {
return nil, trace.Wrap(err)
}
return buf.Bytes(), nil
}
// Packet is a packet interface.
type Packet interface {
// Bytes returns whole packet bytes.
Bytes() []byte
// Data returns packet data without data related to Header.
Data() []byte
// Header returns packet Header definition.
Header() PacketHeader
// Type returns packet type ID.
Type() uint8
}
// BasicPacket implements the Packet interfaces allowing to operate on
// PacketHeader and get underlying packet type.
type BasicPacket struct {
header PacketHeader
data []byte
raw bytes.Buffer
}
// Bytes returns raw packet bytes.
func (g BasicPacket) Bytes() []byte {
return g.raw.Bytes()
}
// Data is the packet data bytes without header.
func (g BasicPacket) Data() []byte {
return g.data
}
// Header is the parsed packet header.
func (g BasicPacket) Header() PacketHeader {
return g.header
}
// Type is the parsed packet header.
func (g BasicPacket) Type() uint8 {
return g.header.Type
}
// ReadPacket reads a single full packet from the reader.
func ReadPacket(r io.Reader) (*BasicPacket, error) {
var buff bytes.Buffer
tr := io.TeeReader(r, &buff)
// Read 8-byte packet header.
var headerBytes [packetHeaderSize]byte
if _, err := io.ReadFull(tr, headerBytes[:]); err != nil {
return nil, trace.ConvertSystemError(err)
}
// Unmarshal packet header from the binary form.
var header PacketHeader
if err := binary.Read(bytes.NewReader(headerBytes[:]), binary.BigEndian, &header); err != nil {
return nil, trace.Wrap(err)
}
// Read packet data. Packet length includes header.
dataBytes := make([]byte, header.Length-packetHeaderSize)
if _, err := io.ReadFull(tr, dataBytes); err != nil {
return nil, trace.ConvertSystemError(err)
}
p := &BasicPacket{
header: header,
data: dataBytes,
raw: buff,
}
return p, nil
}
// NewBasicPacket creates a new BasicPacket instance with the specified
// PacketHeader and data.
func NewBasicPacket(header PacketHeader, data []byte) (*BasicPacket, error) {
headerBytes, err := header.Marshal()
if err != nil {
return nil, trace.Wrap(err)
}
raw := bytes.NewBuffer(append(headerBytes, data...))
return &BasicPacket{
header: header,
data: data,
raw: *raw,
}, nil
}
// ToSQLPacket tries to convert basicPacket to MSServer SQL packet.
func ToSQLPacket(p *BasicPacket) (out Packet, err error) {
defer func() {
if r := recover(); r != nil {
err = trace.BadParameter("failed to convert packet to SQL packet: %v @ %v", r, debug.Stack())
}
}()
switch p.Type() {
case PacketTypeRPCRequest:
sqlBatch, err := toRPCRequest(p)
if err != nil {
return p, trace.Wrap(err)
}
return sqlBatch, trace.Wrap(err)
case PacketTypeSQLBatch:
rpcRequest, err := toSQLBatch(p)
if err != nil {
return p, trace.Wrap(err)
}
return rpcRequest, trace.Wrap(err)
}
return p, trace.Wrap(err)
}
// makePacket prepends header to the provided packet data.
func makePacket(pktType uint8, pktData []byte) ([]byte, error) {
header := PacketHeader{
Type: pktType,
Status: PacketStatusLast,
Length: uint16(packetHeaderSize + len(pktData)),
}
headerBytes, err := header.Marshal()
if err != nil {
return nil, trace.Wrap(err)
}
return append(headerBytes, pktData...), nil
}
// IsFinalPacket returns true there are no more packets on the message.
func IsFinalPacket(packet Packet) bool {
return packet.Header().Status&PacketStatusLast == 1
}
/*
* 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 protocol
import (
"bytes"
"io"
"github.com/gravitational/trace"
mssql "github.com/microsoft/go-mssqldb"
)
// PreLoginPacket represents a Pre-Login packet which is sent by the client
// to set up context for login.
//
// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/60f56408-0188-4cd5-8b90-25c6f2423868
type PreLoginPacket struct {
packet Packet
}
// ReadPreLoginPacket reads Pre-Login packet from the reader.
func ReadPreLoginPacket(r io.Reader) (*PreLoginPacket, error) {
pkt, err := ReadPacket(r)
if err != nil {
return nil, trace.Wrap(err)
}
if pkt.Type() != PacketTypePreLogin {
return nil, trace.BadParameter("expected Pre-Login packet, got: %#v", pkt)
}
return &PreLoginPacket{
packet: pkt,
}, nil
}
// WritePreLoginResponse writes response to the Pre-Login packet to the writer.
func WritePreLoginResponse(w io.Writer) error {
var buf bytes.Buffer
if err := mssql.WritePreLoginFields(&buf, preLoginOptions); err != nil {
return trace.Wrap(err)
}
pkt, err := makePacket(PacketTypeResponse, buf.Bytes())
if err != nil {
return trace.Wrap(err)
}
_, err = w.Write(pkt)
if err != nil {
return trace.Wrap(err)
}
return nil
}
/*
* 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 protocol
import (
"bytes"
"encoding/binary"
"fmt"
"io"
"github.com/gravitational/trace"
mssql "github.com/microsoft/go-mssqldb"
"github.com/microsoft/go-mssqldb/msdsn"
)
// procIDToName maps procID to the special stored procedure name
// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/619c43b6-9495-4a58-9e49-a4950db245b3
var procIDToName = []string{
1: "Sp_Cursor",
2: "Sp_CursorOpen",
3: "Sp_CursorPrepare",
4: "Sp_CursorExecute",
5: "Sp_CursorPrepExec",
6: "Sp_CursorUnprepare",
7: "Sp_CursorFetch",
8: "Sp_CursorOption",
9: "Sp_CursorClose",
10: "Sp_ExecuteSql",
11: "Sp_Prepare",
12: "Sp_Execute",
13: "Sp_PrepExec",
14: "Sp_PrepExecRpc",
15: "Sp_Unprepare",
}
// RPCRequest defines client RPC Request packet:
// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/619c43b6-9495-4a58-9e49-a4950db245b3
type RPCRequest struct {
Packet
// ProcName contains name of the procedure to be executed.
ProcName string
// Parameters contains list of RPC parameters.
Parameters []string
}
func toRPCRequest(p Packet) (*RPCRequest, error) {
if p.Type() != PacketTypeRPCRequest {
return nil, trace.BadParameter("expected SQLBatch packet, got: %#v", p.Type())
}
data := p.Data()
r := bytes.NewReader(p.Data())
var headersLength uint32
if err := binary.Read(r, binary.LittleEndian, &headersLength); err != nil {
return nil, trace.Wrap(err)
}
if _, err := r.Seek(int64(headersLength), io.SeekStart); err != nil {
return nil, trace.ConvertSystemError(err)
}
var length uint16
if err := binary.Read(r, binary.LittleEndian, &length); err != nil {
return nil, trace.Wrap(err)
}
var procName string
var err error
// If the first USHORT contains 0xFFFF the following USHORT contains the PROCID.
// Otherwise, NameLenProcID contains the parameter name length and parameter name.
if length == procIDSwitchRPCRequest {
var procID uint16
if err := binary.Read(r, binary.LittleEndian, &procID); err != nil {
return nil, trace.Wrap(err)
}
procName, err = getProcName(procID)
if err != nil {
return nil, trace.BadParameter("failed to get procedure name")
}
} else {
procName, err = readUcs2(r, 2*int(length))
if err != nil {
return nil, trace.Wrap(err)
}
}
var flags uint16
if err := binary.Read(r, binary.LittleEndian, &flags); err != nil {
return nil, trace.Wrap(err)
}
// offset the reader by 2 bytes.
if _, err := r.Seek(2, io.SeekCurrent); err != nil {
return nil, trace.ConvertSystemError(err)
}
tds := mssql.NewTdsBuffer(data[int(r.Size())-r.Len():], r.Len())
typeId, err := tds.ReadByte()
if err != nil {
return nil, trace.Wrap(err)
}
// pass nil for crypto parameter, we are dealing with unencrypted data here.
ti := mssql.ReadTypeInfo(tds, typeId, nil, msdsn.EncodeParameters{GuidConversion: false})
val := ti.Reader(&ti, tds, nil)
return &RPCRequest{
Packet: p,
ProcName: procName,
Parameters: getParameters(val),
}, nil
}
func getParameters(val any) []string {
if val == nil {
return nil
}
return []string{fmt.Sprintf("%v", val)}
}
func getProcName(procID uint16) (string, error) {
if int(procID) >= len(procIDToName) {
return "unknownProc", nil
}
var procName string
if procName = procIDToName[procID]; procName == "" {
return "", trace.BadParameter("unmapped procID")
}
return procName, nil
}
/*
* 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 protocol
import (
"bytes"
"encoding/binary"
"io"
"github.com/gravitational/trace"
mssql "github.com/microsoft/go-mssqldb"
)
// SQLBatch is a representation of MSServer SQL Batch packet.
// https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-tds/f2026cd3-9a46-4a3f-9a08-f63140bcbbe3
type SQLBatch struct {
Packet
// SQLText contains text batch query.
SQLText string
}
func toSQLBatch(p Packet) (*SQLBatch, error) {
if p.Type() != PacketTypeSQLBatch {
return nil, trace.BadParameter("expected SQLBatch packet, got: %v", p.Type())
}
r := bytes.NewReader(p.Data())
var headersLength uint32
// The packet header if present only in the first packet.
if int(p.Header().PacketID) == 1 {
if err := binary.Read(r, binary.LittleEndian, &headersLength); err != nil {
return nil, trace.Wrap(err)
}
}
if _, err := r.Seek(int64(headersLength), io.SeekStart); err != nil {
return nil, trace.ConvertSystemError(err)
}
s, err := mssql.ParseUCS2String(p.Data()[r.Size()-int64(r.Len()):])
if err != nil {
return nil, trace.Wrap(err)
}
return &SQLBatch{
Packet: p,
SQLText: s,
}, nil
}
/*
* 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 protocol
import (
"io"
"github.com/gravitational/trace"
mssql "github.com/microsoft/go-mssqldb"
)
// WriteStreamResponse writes stream response packet to the writer.
func WriteStreamResponse(w io.Writer, tokens []mssql.Token) error {
var data []byte
for _, token := range tokens {
bytes, err := token.Marshal()
if err != nil {
return trace.Wrap(err)
}
data = append(data, bytes...)
}
pkt, err := makePacket(PacketTypeResponse, data)
if err != nil {
return trace.Wrap(err)
}
_, err = w.Write(pkt)
if err != nil {
return trace.Wrap(err)
}
return nil
}
// WriteErrorResponse writes error response to the client.
func WriteErrorResponse(w io.Writer, err error) error {
return WriteStreamResponse(w, []mssql.Token{
&mssql.Error{
Number: errorNumber,
Class: errorClassSecurity,
Message: err.Error(),
},
mssql.DoneToken(),
})
}
/*
* 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 protocol
import (
"io"
mssql "github.com/microsoft/go-mssqldb"
)
func readUcs2(r io.Reader, numchars int) (string, error) {
buf := make([]byte, numchars)
_, err := io.ReadFull(r, buf)
if err != nil {
return "", err
}
return mssql.ParseUCS2String(buf)
}
/*
* 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 tdp
import (
"bufio"
"errors"
"io"
"net"
"slices"
"sync"
"github.com/gravitational/trace"
)
const (
// MaxAlertMessageLength is somewhat arbitrary, as it is only sent *to*
// the browser (Teleport never receives this message, so won't be decoding it)
MaxAlertMessageLength = 10240
// MaxPathLength is somewhat arbitrary because we weren't able to determine
// a precise value to set it to: https://github.com/gravitational/teleport/issues/14950#issuecomment-1341632465
// The limit is kept as an additional defense-in-depth measure.
MaxPathLength = 10240
MaxClipboardDataLength = 1024 * 1024 // 1MB
MaxFileReadWriteLength = 1024 * 1024 // 1MB
)
var (
ClipDataMaxLenErr = trace.LimitExceeded("clipboard sync failed: clipboard data exceeded maximum length")
StringMaxLenErr = trace.LimitExceeded("TDP string length exceeds allowable limit")
FileReadWriteMaxLenErr = trace.LimitExceeded("TDP file read or write message exceeds maximum size limit")
MFADataMaxLenErr = trace.LimitExceeded("MFA challenge data exceeds maximum length")
)
// IsNonFatalErr returns whether or not an error arising from
// the tdp package should be interpreted as fatal or non-fatal
// for an ongoing TDP connection.
func IsNonFatalErr(err error) bool {
if err == nil {
return false
}
return errors.Is(err, ClipDataMaxLenErr) ||
errors.Is(err, StringMaxLenErr) ||
errors.Is(err, FileReadWriteMaxLenErr) ||
errors.Is(err, MFADataMaxLenErr)
}
// IsFatalErr returns the inverse of IsNonFatalErr
// (except for if err == nil, for which both functions return false)
func IsFatalErr(err error) bool {
if err == nil {
return false
}
return !IsNonFatalErr(err)
}
type Message interface {
Encode() ([]byte, error)
}
type MessageReader interface {
ReadMessage() (Message, error)
}
type MessageWriter interface {
WriteMessage(Message) error
}
type MessageReadWriter interface {
MessageReader
MessageWriter
}
type MessageReadWriteCloser interface {
MessageReadWriter
Close() error
}
// Conn is a desktop protocol connection.
// It converts between a stream of bytes (io.ReadWriter) and a stream of
// Teleport Desktop Protocol (TDP) messages.
type Conn struct {
rwc io.ReadWriteCloser
writeMu sync.Mutex
bufr *bufio.Reader
decode Decoder
closeOnce sync.Once
constructWarning WarningConstructor
// localAddr and remoteAddr will be set if rw is
// a conn that provides these fields
localAddr net.Addr
remoteAddr net.Addr
}
type ByteReader interface {
io.Reader
io.ByteReader
}
// Decoder is a function that decodes incoming data
// into a Message.
type Decoder func(ByteReader) (Message, error)
// DecoderAdapter adapts a Decoder that works with a standard io.Reader
func DecoderAdapter(f func(io.Reader) (Message, error)) Decoder {
return func(br ByteReader) (Message, error) {
return f(br)
}
}
// NewConn creates a new Conn on top of a ReadWriter, for example a TCP
// connection. If the provided ReadWriter also implements srv.TrackingConn,
// then its LocalAddr() and RemoteAddr() will apply to this Conn.
func NewConn(rwc io.ReadWriteCloser, decoder Decoder, wc WarningConstructor) *Conn {
br := bufio.NewReader(rwc)
c := &Conn{
rwc: rwc,
bufr: br,
decode: decoder,
constructWarning: wc,
}
if tc, ok := rwc.(srvTrackingConn); ok {
c.localAddr = tc.LocalAddr()
c.remoteAddr = tc.RemoteAddr()
}
return c
}
// srvTrackingConn should be kept in sync with
// lib/srv.TrackingConn. It is duplicated here
// to avoid placing a dependency on the lib/srv
// package, which is incompatible with Windows.
type srvTrackingConn interface {
LocalAddr() net.Addr
RemoteAddr() net.Addr
Close() error
}
// Close closes the connection if the underlying reader can be closed.
func (c *Conn) Close() error {
var err error
c.closeOnce.Do(func() {
err = c.rwc.Close()
})
return err
}
// PeekNextByte peeks at the next byte without consuming it.
func (c *Conn) PeekNextByte() (byte, error) {
b, err := c.bufr.ReadByte()
if err != nil {
return 0, trace.Wrap(err)
}
if err := c.bufr.UnreadByte(); err != nil {
return 0, trace.Wrap(err)
}
return b, nil
}
// WarningConstructor is a function that constructs a TDP or TDPB
// alert message with warning severity to be sent to the client.
type WarningConstructor func(string) Message
func (c *Conn) sendWarning(warning string) error {
return c.WriteMessage(c.constructWarning(warning))
}
// ReadMessage reads the next incoming message from the connection.
func (c *Conn) ReadMessage() (Message, error) {
for {
m, err := c.decode(c.bufr)
if err != nil && IsNonFatalErr(err) {
if warnError := c.sendWarning(err.Error()); warnError != nil {
return nil, trace.Wrap(warnError, "error sending alert message in response to decode error: %v", err)
}
// Warning sent. Try reading the next message
continue
}
// err may still be non-nil
return m, trace.Wrap(err)
}
}
// WriteMessage sends a message to the connection.
func (c *Conn) WriteMessage(m Message) error {
buf, err := m.Encode()
if err != nil {
return trace.Wrap(err)
}
c.writeMu.Lock()
_, err = c.rwc.Write(buf)
c.writeMu.Unlock()
return trace.Wrap(err)
}
// LocalAddr returns local address
func (c *Conn) LocalAddr() net.Addr {
return c.localAddr
}
// RemoteAddr returns remote address
func (c *Conn) RemoteAddr() net.Addr {
return c.remoteAddr
}
// Interceptor intercepts messages. It should return
// the [potentially modified] message(s) in order to pass it on to the
// other end of the connection, or it may swallow the message by returning
// a nil or empty slice. Returned slices should not contain nil messages.
type Interceptor func(message Message) ([]Message, error)
// ReadWriteInterceptor wraps an existing 'MessageReadWriteCloser' and runs the
// provided interceptor functions in the read and/or write paths. Allows callers
// to snoop and modify messages as they pass through the 'MessageReadWriteCloser'.
type ReadWriteInterceptor struct {
// The underlying read/writer to intercept messages on
src MessageReadWriteCloser
// The interceptor to run in the write path
writeInterceptor Interceptor
// A closure over the interceptor to run in the read path
readAdapter func() (Message, error)
}
// Message slices returned by interceptor functions should not include
// nil messages, but a little defensive programming can prevent a crash.
// 'removeNilMessages' will return a subslice of the input slice with
// nil messages removed.
func removeNilMessages(msgs []Message) []Message {
return slices.DeleteFunc(msgs, func(msg Message) bool {
return msg == nil
})
}
// readInterceptorAdapter closes over an internal slice that keeps track of
// of messages returned by the read interceptor callback (if present).
func readInterceptorAdapter(src MessageReader, i Interceptor) func() (Message, error) {
if i == nil {
return src.ReadMessage
}
var msgs []Message
return func() (Message, error) {
// len(msgs) == 0 - initial case / empty cache
// len(msgs) > 0 - Return a cached message
for len(msgs) == 0 {
// Try reading a message
m, err := src.ReadMessage()
if err != nil {
return nil, err
}
// The interceptor may return an empty slice and nil error
// In that case, we'll try again via the loop.
msgs, err = i(m)
if err != nil {
return nil, err
}
msgs = removeNilMessages(msgs)
}
msg := msgs[0]
msgs = msgs[1:]
return msg, nil
}
}
// NewReadWriteInterceptor creates a new 'ReadWriteInterceptor' that intercepts messages on 'src'.
// The provided interceptor callbacks may be nil.
func NewReadWriteInterceptor(src MessageReadWriteCloser, readInterceptor, writeInterceptor Interceptor) *ReadWriteInterceptor {
return &ReadWriteInterceptor{
src: src,
writeInterceptor: writeInterceptor,
readAdapter: readInterceptorAdapter(src, readInterceptor),
}
}
// WriteMessage passes the message to the write interceptor (if provided)
// for omition or modification before writing the message to the underlying
// writer.
func (i *ReadWriteInterceptor) WriteMessage(m Message) error {
if i.writeInterceptor == nil {
// No interceptor found
return i.src.WriteMessage(m)
}
out, err := i.writeInterceptor(m)
if err != nil {
return err
}
for _, msg := range out {
if msg == nil {
continue
}
if err = i.src.WriteMessage(msg); err != nil {
return err
}
}
return nil
}
// ReadMessage reads from the underlying reader and passes them to the
// read interceptor (if provided) for omition or modification before
// returning the next message.
func (i *ReadWriteInterceptor) ReadMessage() (Message, error) {
return i.readAdapter()
}
// Close closes the underlying 'MessageReadWriteCloser'
func (i *ReadWriteInterceptor) Close() error {
return i.src.Close()
}
// copyMessages behaves similarly to io.Copy except it deals with Message types.
// It reads messages from 'src' and writes them to 'dst' until an error is received.
// It does *not* forward an EOF received from the reader, but returns nil in the happy path.
func copyMessages(dst MessageWriter, src MessageReader) error {
for {
msg, err := src.ReadMessage()
if errors.Is(err, io.EOF) {
return nil
} else if err != nil {
return err
}
if err := dst.WriteMessage(msg); err != nil {
return err
}
}
}
// ConnProxy handles bi-directional copying of messages from server <-> client.
type ConnProxy struct {
server MessageReadWriteCloser
client MessageReadWriteCloser
}
// NewConnProxy returns a new ConnProxy.
func NewConnProxy(client, server MessageReadWriteCloser) ConnProxy {
return ConnProxy{
server: server,
client: client,
}
}
// Run handles bi-directional copying of messages from server <-> client until
// an IO error occurs (or EOF is received from either side). It always calls
// 'close' on both streams before exiting and returns any errors occurred from
// reading, writing, or closing both streams.
func (c *ConnProxy) Run() error {
wg := sync.WaitGroup{}
// Copy in both directions
var clientToServerErr, serverToClientErr error
wg.Go(func() {
err := copyMessages(c.client, c.server)
serverToClientErr = trace.NewAggregate(err, c.client.Close())
})
wg.Go(func() {
err := copyMessages(c.server, c.client)
clientToServerErr = trace.NewAggregate(err, c.server.Close())
})
wg.Wait()
return trace.NewAggregate(clientToServerErr, serverToClientErr)
}
// EncodeTo calls 'Encode' on the given message and writes it to 'w'.
func EncodeTo(w io.Writer, msg Message) error {
data, err := msg.Encode()
if err != nil {
return trace.Wrap(err)
}
_, err = w.Write(data)
return trace.Wrap(err)
}
// Teleport
// Copyright (C) 2026 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 tdp
import (
"context"
"errors"
"log/slog"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/mfa"
)
var (
ErrUnexpectedTDPMessageType = errors.New("unexpected message type")
)
// convertChallenge converts an MFA challenge to a Message. Returns
// a non-nil error if the conversion fails
type convertChallenge func(*proto.MFAAuthenticateChallenge) (Message, error)
// asMFAResponse returns:
// - ErrUnexpectedTDPMessageType if a valid messages was received but was not an MFA message.
// - Any other non-nil error if there was an error interpreting the message.
// - nil if a valid, non-nil MFA messages was found.
type asMFAResponse func(Message) (*proto.MFAAuthenticateResponse, error)
// NewMfaPrompt constructs a function that reads, encodes, and sends an MFA challenge to the client,
// then waits for the corresponding MFA response message. It caches any non-MFA messages received so
// that they may be forwarded to the server later on.
func NewMFAPrompt(rw MessageReadWriter, asResponse asMFAResponse, toMessage convertChallenge, withheld *[]Message, log *slog.Logger) mfa.PromptFunc {
return func(ctx context.Context, chal *proto.MFAAuthenticateChallenge) (*proto.MFAAuthenticateResponse, error) {
challengeMsg, err := toMessage(chal)
if err != nil {
return nil, trace.Wrap(err)
}
log.DebugContext(ctx, "Writing MFA challenge to client")
if err = rw.WriteMessage(challengeMsg); err != nil {
return nil, trace.Wrap(err)
}
for {
msg, err := rw.ReadMessage()
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := asResponse(msg)
if err != nil {
if errors.Is(err, ErrUnexpectedTDPMessageType) {
// Withhold this non-MFA message and try reading again
log.DebugContext(ctx, "Received non-MFA message", "message", msg)
*withheld = append(*withheld, msg)
continue
} else {
log.DebugContext(ctx, "Error receiving MFA response", "error", err)
// Unexpected error occurred while inspecting the message
return nil, trace.Wrap(err)
}
}
// Found our MFA response!
log.DebugContext(ctx, "Received MFA response")
return resp, nil
}
}
}
/*
* 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 tdp
import "image/png"
// PNGEncoder returns the encoder used for PNG Frames.
// It is not safe for concurrent use.
func PNGEncoder() *png.Encoder {
return &png.Encoder{
CompressionLevel: png.BestSpeed,
BufferPool: &pool{},
}
}
// pool implements png.EncoderBufferPool,
// allowing us to reuse encoding resources
type pool struct {
b *png.EncoderBuffer
}
// all encoding happens in a single thread, so we don't
// need anything as sophisticated as a sync.Pool here
func (p *pool) Get() *png.EncoderBuffer { return p.b }
func (p *pool) Put(eb *png.EncoderBuffer) { p.b = eb }
/*
* 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 regular
import (
"context"
"fmt"
"log/slog"
"net"
"strings"
"github.com/gravitational/trace"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/agentless"
"github.com/gravitational/teleport/lib/proxy"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/sshagent"
"github.com/gravitational/teleport/lib/utils"
)
// PROXYHeaderSigner allows to sign PROXY headers for securely propagating original client IP information
type PROXYHeaderSigner interface {
SignPROXYHeader(source, destination net.Addr) ([]byte, error)
}
// CertAuthorityGetter allows to get cluster's host CA for verification of signed PROXY headers.
// We define our own version to avoid circular dependencies in multiplexer package (it can't depend on 'services'),
// where this function is used.
type CertAuthorityGetter = func(ctx context.Context, id types.CertAuthID, loadKeys bool) (types.CertAuthority, error)
// proxySubsys implements an SSH subsystem for proxying listening sockets from
// remote hosts to a proxy client (AKA port mapping)
type proxySubsys struct {
proxySubsysRequest
router *proxy.Router
ctx *srv.ServerContext
logger *slog.Logger
closeC chan error
proxySigner PROXYHeaderSigner
localCluster string
}
// parseProxySubsys looks at the requested subsystem name and returns a fully configured
// proxy subsystem
//
// proxy subsystem name can take the following forms:
//
// "proxy:host:22" - standard SSH request to connect to host:22 on the 1st cluster
// "proxy:@clustername" - Teleport request to connect to an auth server for cluster with name 'clustername'
// "proxy:host:22@clustername" - Teleport request to connect to host:22 on cluster 'clustername'
// "proxy:host:22@namespace@clustername"
func (s *Server) parseProxySubsysRequest(ctx context.Context, request string) (proxySubsysRequest, error) {
s.logger.DebugContext(ctx, "parsing proxy subsystem request", "request", request)
var (
clusterName string
targetHost string
targetPort string
paramMessage = fmt.Sprintf("invalid format for proxy request: %q, expected 'proxy:host:port@cluster'", request)
)
const prefix = "proxy:"
// get rid of 'proxy:' prefix:
if strings.Index(request, prefix) != 0 {
return proxySubsysRequest{}, trace.BadParameter("%s", paramMessage)
}
requestBody := strings.TrimPrefix(request, prefix)
namespace := apidefaults.Namespace
parts := strings.Split(requestBody, "@")
var err error
switch {
case len(parts) == 0: // "proxy:"
return proxySubsysRequest{}, trace.BadParameter("%s", paramMessage)
case len(parts) == 1: // "proxy:host:22"
targetHost, targetPort, err = utils.SplitHostPort(parts[0])
if err != nil {
return proxySubsysRequest{}, trace.BadParameter("%s", paramMessage)
}
case len(parts) == 2: // "proxy:@clustername" or "proxy:host:22@clustername"
if parts[0] != "" {
targetHost, targetPort, err = utils.SplitHostPort(parts[0])
if err != nil {
return proxySubsysRequest{}, trace.BadParameter("%s", paramMessage)
}
}
clusterName = parts[1]
if clusterName == "" && targetHost == "" {
return proxySubsysRequest{}, trace.BadParameter("invalid format for proxy request: missing cluster name or target host in %q", request)
}
case len(parts) >= 3: // "proxy:host:22@namespace@clustername"
clusterName = strings.Join(parts[2:], "@")
namespace = parts[1]
targetHost, targetPort, err = utils.SplitHostPort(parts[0])
if err != nil {
return proxySubsysRequest{}, trace.BadParameter("%s", paramMessage)
}
}
return proxySubsysRequest{
namespace: namespace,
host: targetHost,
port: targetPort,
clusterName: clusterName,
}, nil
}
// parseProxySubsys decodes a proxy subsystem request and sets up a proxy subsystem instance.
// See parseProxySubsysRequest for details on the request format.
func (s *Server) parseProxySubsys(ctx context.Context, request string, serverContext *srv.ServerContext) (*proxySubsys, error) {
req, err := s.parseProxySubsysRequest(ctx, request)
if err != nil {
return nil, trace.Wrap(err)
}
subsys, err := newProxySubsys(ctx, serverContext, s, req)
if err != nil {
return nil, trace.Wrap(err)
}
return subsys, nil
}
// proxySubsysRequest encodes proxy subsystem request parameters.
type proxySubsysRequest struct {
namespace string
host string
port string
clusterName string
}
func (p *proxySubsysRequest) String() string {
return fmt.Sprintf("host=%v, port=%v, cluster=%v", p.host, p.port, p.clusterName)
}
// SpecifiedPort returns whether the port is set, and it has a non-zero value
func (p *proxySubsysRequest) SpecifiedPort() bool {
return len(p.port) > 0 && p.port != "0"
}
// SetDefaults sets default values.
func (p *proxySubsysRequest) SetDefaults() {
if p.namespace == "" {
p.namespace = apidefaults.Namespace
}
}
// newProxySubsys is a helper that creates a proxy subsystem from
// a port forwarding request, used to implement ProxyJump feature in proxy
// and reuse the code
func newProxySubsys(ctx context.Context, serverContext *srv.ServerContext, srv *Server, req proxySubsysRequest) (*proxySubsys, error) {
req.SetDefaults()
if req.clusterName == "" && serverContext.Identity.RouteToCluster != "" {
srv.logger.DebugContext(ctx, "Proxy subsystem: routing user to cluster based on the route to cluster extension",
"user", serverContext.Identity.TeleportUser,
"cluster", serverContext.Identity.RouteToCluster,
)
req.clusterName = serverContext.Identity.RouteToCluster
}
if req.clusterName != "" && srv.proxyClusterGetter != nil {
checker, err := srv.clusterGetterWithAccessChecker(serverContext)
if err != nil {
return nil, trace.Wrap(err)
}
if _, err := checker.Cluster(ctx, req.clusterName); err != nil {
return nil, trace.BadParameter("invalid format for proxy request: unknown cluster %q", req.clusterName)
}
}
srv.logger.DebugContext(ctx, "successfully created proxy subsystem request", "request", &req)
return &proxySubsys{
proxySubsysRequest: req,
ctx: serverContext,
logger: slog.With(teleport.ComponentKey, teleport.ComponentSubsystemProxy),
closeC: make(chan error),
router: srv.router,
proxySigner: srv.proxySigner,
localCluster: serverContext.ClusterName,
}, nil
}
func (t *proxySubsys) String() string {
return fmt.Sprintf("proxySubsys(cluster=%s/%s, host=%s, port=%s)",
t.namespace, t.clusterName, t.host, t.port)
}
// Start is called by Golang's ssh when it needs to engage this subsystem (typically to establish
// a mapping connection between a client & remote node we're proxying to)
func (t *proxySubsys) Start(ctx context.Context, sconn *ssh.ServerConn, ch ssh.Channel, req *ssh.Request, serverContext *srv.ServerContext) error {
// once we start the connection, update logger to include component fields
t.logger = t.logger.With(
"src", sconn.RemoteAddr().String(),
"dst", sconn.LocalAddr().String(),
"subsystem", t.String(),
)
t.logger.DebugContext(ctx, "Starting subsystem")
clientAddr := sconn.RemoteAddr()
// connect to a site's auth server
if t.host == "" {
return t.proxyToSite(ctx, ch, t.clusterName, clientAddr, sconn.LocalAddr())
}
// connect to a server
return t.proxyToHost(ctx, ch, clientAddr, sconn.LocalAddr())
}
// proxyToSite establishes a proxy connection from the connected SSH client to the
// auth server of the requested remote site
func (t *proxySubsys) proxyToSite(ctx context.Context, ch ssh.Channel, clusterName string, clientSrcAddr, clientDstAddr net.Addr) error {
t.logger.DebugContext(ctx, "attempting to proxy connection to auth server", "local_cluster", t.localCluster, "proxied_cluster", clusterName)
conn, err := t.router.DialSite(ctx, clusterName, clientSrcAddr, clientDstAddr)
if err != nil {
return trace.Wrap(err)
}
t.logger.InfoContext(ctx, "Connected to cluster", "cluster", clusterName, "address", conn.RemoteAddr())
go func() {
t.close(utils.ProxyConn(ctx, ch, conn))
}()
return nil
}
// proxyToHost establishes a proxy connection from the connected SSH client to the
// requested remote node (t.host:t.port) via the given site
func (t *proxySubsys) proxyToHost(ctx context.Context, ch ssh.Channel, clientSrcAddr, clientDstAddr net.Addr) error {
t.logger.DebugContext(ctx, "proxying connection to target host", "host", t.host, "port", t.port, "exact_port", t.SpecifiedPort())
authClient, err := t.router.GetSiteClient(ctx, t.localCluster)
if err != nil {
return trace.Wrap(err)
}
certGen, err := t.router.GetSiteClient(ctx, t.clusterName)
if err != nil {
return trace.Wrap(err)
}
identity := t.ctx.Identity
signer := agentless.SignerFromSSHIdentity(identity.UnmappedIdentity, authClient, certGen, t.clusterName, identity.TeleportUser)
aGetter := func() (sshagent.Client, error) {
return t.ctx.StartAgentChannel()
}
conn, err := t.router.DialHost(ctx, identity.UnmappedIdentity.ScopePin, clientSrcAddr, clientDstAddr, t.host, t.port, t.clusterName, t.ctx.Identity.UnstableClusterAccessChecker, aGetter, signer)
if err != nil {
return trace.Wrap(err)
}
go func() {
t.close(utils.ProxyConn(ctx, ch, conn))
}()
return nil
}
func (t *proxySubsys) close(err error) {
t.closeC <- err
}
func (t *proxySubsys) Wait() error {
return <-t.closeC
}
/*
* 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 regular
import (
"context"
"encoding/json"
"errors"
"io"
"log/slog"
"os"
"sync"
"time"
"github.com/gravitational/trace"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/srv"
reexecutils "github.com/gravitational/teleport/lib/sshutils/reexec"
sftputils "github.com/gravitational/teleport/lib/sshutils/sftp"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
"github.com/gravitational/teleport/session/reexec/reexecsftp"
sessionsftputils "github.com/gravitational/teleport/session/sftputils"
)
type sftpSubsys struct {
logger *slog.Logger
fileTransferReq *reexecsftp.FileTransferRequest
sftpCmd *reexec.CommandExecutor
serverCtx *srv.ServerContext
// waitForOutputStreams tracks goroutines that copy stderr/stdout from child
// reexec and shell processes. This is necessary due to the use of custom pipes,
// which exec.Cmd does not wait for closure of in cmd.Wait().
waitForOutputStreams sync.WaitGroup
}
func newSFTPSubsys(fileTransferReq *reexecsftp.FileTransferRequest) (*sftpSubsys, error) {
return &sftpSubsys{
logger: slog.With(teleport.ComponentKey, "subsystem:sftp"),
fileTransferReq: fileTransferReq,
}, nil
}
func (s *sftpSubsys) Start(ctx context.Context,
serverConn *ssh.ServerConn,
ch ssh.Channel, req *ssh.Request,
serverCtx *srv.ServerContext,
) error {
// Check that file copying is allowed Node-wide again here in case
// this connection was proxied, the proxy doesn't know if file copying
// is allowed for certain Nodes.
if !serverCtx.AllowFileCopying {
serverCtx.GetServer().EmitAuditEvent(context.WithoutCancel(ctx), &apievents.SFTP{
Metadata: apievents.Metadata{
Code: events.SFTPDisallowedCode,
Type: events.SFTPEvent,
Time: time.Now(),
},
UserMetadata: serverCtx.Identity.GetUserMetadata(),
ServerMetadata: serverCtx.GetServer().EventMetadata(),
Error: srv.ErrNodeFileCopyingNotPermitted.Error(),
})
return srv.ErrNodeFileCopyingNotPermitted
}
s.serverCtx = serverCtx
// Create two sets of anonymous pipes to give the child process
// access to the SSH channel
chReadPipeOut, chReadPipeIn, err := os.Pipe()
if err != nil {
return trace.Wrap(err)
}
defer chReadPipeOut.Close()
chWritePipeOut, chWritePipeIn, err := os.Pipe()
if err != nil {
return trace.Wrap(err)
}
defer chWritePipeIn.Close()
// Create anonymous pipe that the child will send audit information
// over
auditPipeOut, auditPipeIn, err := os.Pipe()
if err != nil {
return trace.Wrap(err)
}
defer auditPipeIn.Close()
// Create child process to handle SFTP connection
execRequest, err := srv.NewExecRequest(serverCtx, reexecconstants.SFTPSubCommand)
if err != nil {
return trace.Wrap(err)
}
if err := serverCtx.SetExecRequest(execRequest); err != nil {
return trace.Wrap(err)
}
if err := serverCtx.SetSSHRequest(req); err != nil {
return trace.Wrap(err)
}
s.sftpCmd, err = serverCtx.ConfigureCommand(map[reexec.FileFD]*os.File{
reexec.StdinFile: chReadPipeOut,
reexec.StdoutFile: chWritePipeIn,
reexec.StderrFile: auditPipeIn,
})
if err != nil {
return trace.Wrap(err)
}
// Capture stderr.
stderrR, stderrW, err := os.Pipe()
if err != nil {
return trace.Wrap(err)
}
defer stderrW.Close()
s.sftpCmd.Stderr = stderrW
s.waitForOutputStreams.Go(func() {
defer stderrR.Close()
childErr, err := reexecutils.ReadChildErrorWithContext(stderrR, &reexecutils.ErrorContext{
DecisionContext: s.serverCtx.Identity.AccessPermit.GetDecisionContext(),
Login: s.serverCtx.Identity.Login,
})
if err != nil {
s.logger.WarnContext(context.WithoutCancel(ctx), "Failed to read child process stderr", "error", err)
return
}
if childErr == "" {
return
}
if _, err := io.WriteString(ch.Stderr(), childErr); err != nil {
s.logger.WarnContext(context.WithoutCancel(ctx), "Failed to propagate child process stderr to client", "error", err)
}
})
s.logger.DebugContext(ctx, "starting SFTP process")
err = s.sftpCmd.Start()
if err != nil {
return trace.Wrap(err)
}
if err := s.sftpCmd.Continue(); err != nil {
return trace.Wrap(err)
}
// Send the file transfer request if applicable. The SFTP process
// expects the file transfer request data will end with a null byte,
// so if there is no request to send just send a null byte so the
// SFTP process can detect that no request was sent.
encodedReq := []byte{0x0}
if s.fileTransferReq != nil {
encodedReq, err = json.Marshal(s.fileTransferReq)
if err != nil {
return trace.Wrap(err)
}
encodedReq = append(encodedReq, 0x0)
}
_, err = chReadPipeIn.Write(encodedReq)
if err != nil {
return trace.Wrap(err)
}
// Copy the SSH channel to and from the anonymous pipes. The input copy from
// the SSH channel must not gate Wait(), or early child-process failures can
// deadlock waiting for the client to close the channel before we send the
// exit status.
go func() {
defer chReadPipeIn.Close()
if _, err := io.Copy(chReadPipeIn, ch); err != nil && !utils.IsOKNetworkError(err) {
s.logger.WarnContext(ctx, "Failure reading from SFTP subsystem", "error", err)
}
}()
s.waitForOutputStreams.Go(func() {
defer chWritePipeOut.Close()
if _, err := io.Copy(ch, chWritePipeOut); err != nil && !utils.IsOKNetworkError(err) {
s.logger.WarnContext(ctx, "Failure writing to SFTP subsystem", "error", err)
}
})
// Read and emit audit events from the child process
go func() {
defer auditPipeOut.Close()
// Create common fields for events
serverMeta := serverCtx.GetServer().EventMetadata()
sessionMeta := serverCtx.GetSessionMetadata()
userMeta := serverCtx.Identity.GetUserMetadata()
connectionMeta := apievents.ConnectionMetadata{
RemoteAddr: serverConn.RemoteAddr().String(),
LocalAddr: serverConn.LocalAddr().String(),
}
dec := json.NewDecoder(auditPipeOut)
for {
var ev sessionsftputils.Event
if err := dec.Decode(&ev); err != nil {
if !errors.Is(err, io.EOF) {
s.logger.WarnContext(ctx, "Failed to read SFTP event", "error", err)
}
return
}
var event apievents.AuditEvent
if ev.SFTP != nil {
e, err := sftputils.SFTPEventToProto(ev.SFTP)
if err != nil {
s.logger.WarnContext(ctx, "Failed to convert SFTP event", "error", err)
continue
}
e.SetClusterName(serverCtx.ClusterName)
e.ServerMetadata = serverMeta
e.SessionMetadata = sessionMeta
e.UserMetadata = userMeta
e.ConnectionMetadata = connectionMeta
event = e
} else if ev.Summary != nil {
e := sftputils.SFTPSummaryEventToProto(ev.Summary)
e.SetClusterName(serverCtx.ClusterName)
e.ServerMetadata = serverMeta
e.SessionMetadata = sessionMeta
e.UserMetadata = userMeta
e.ConnectionMetadata = connectionMeta
event = e
} else {
s.logger.WarnContext(ctx, "Unknown event type received from SFTP server process")
continue
}
if err := serverCtx.GetServer().EmitAuditEvent(ctx, event); err != nil {
s.logger.WarnContext(ctx, "Failed to emit SFTP event", "error", err)
}
}
}()
return nil
}
func (s *sftpSubsys) Wait() error {
ctx := context.Background()
waitErr := s.sftpCmd.Wait()
s.waitForOutputStreams.Wait()
s.logger.DebugContext(ctx, "SFTP process finished")
s.serverCtx.SendExecResult(ctx, srv.ExecResult{
Command: s.sftpCmd.String(),
Code: s.sftpCmd.ProcessState.ExitCode(),
})
return trace.Wrap(waitErr)
}
/*
* 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 regular
import (
"context"
"encoding/json"
"github.com/gravitational/trace"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/srv"
)
// proxySubsys is an SSH subsystem for easy proxyneling through proxy server
// This subsystem creates a new TCP connection and connects ssh channel
// with this connection
type proxySitesSubsys struct {
srv *Server
}
func parseProxySitesSubsys(name string, srv *Server) (*proxySitesSubsys, error) {
return &proxySitesSubsys{
srv: srv,
}, nil
}
func (t *proxySitesSubsys) String() string {
return "proxySites()"
}
func (t *proxySitesSubsys) Wait() error {
return nil
}
// Start serves a request for "proxysites" custom SSH subsystem. It builds an array of
// service.Site structures, and writes it serialized as JSON back to the SSH client
func (t *proxySitesSubsys) Start(ctx context.Context, sconn *ssh.ServerConn, ch ssh.Channel, req *ssh.Request, serverContext *srv.ServerContext) error {
t.srv.logger.DebugContext(ctx, "starting proxysites subsystem", "server_context", serverContext)
checker, err := t.srv.clusterGetterWithAccessChecker(serverContext)
if err != nil {
return trace.Wrap(err)
}
clusters, err := checker.Clusters(ctx)
if err != nil {
return trace.Wrap(err)
}
// build an arary of services.Site structures:
retval := make([]types.Site, 0, len(clusters))
for _, s := range clusters {
retval = append(retval, types.Site{
Name: s.GetName(),
Status: s.GetStatus(),
LastConnected: s.GetLastConnected(),
})
}
// serialize them into JSON and write back:
data, err := json.Marshal(retval)
if err != nil {
return trace.Wrap(err)
}
if _, err := ch.Write(data); err != nil {
return trace.Wrap(err)
}
return nil
}
/*
* 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 regular implements SSH server that supports multiplexing
// tunneling, SSH connections proxying and only supports Key based auth
package regular
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"maps"
"net"
"os"
"os/user"
"runtime"
"strings"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
semconv "go.opentelemetry.io/otel/semconv/v1.10.0"
oteltrace "go.opentelemetry.io/otel/trace"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
decisionpb "github.com/gravitational/teleport/api/gen/proto/go/teleport/decision/v1alpha1"
stableunixusersv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/stableunixusers/v1"
"github.com/gravitational/teleport/api/observability/tracing"
tracessh "github.com/gravitational/teleport/api/observability/tracing/ssh"
"github.com/gravitational/teleport/api/types"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/bpf"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/componentfeatures"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/inventory"
"github.com/gravitational/teleport/lib/labels"
"github.com/gravitational/teleport/lib/limiter"
"github.com/gravitational/teleport/lib/proxy"
"github.com/gravitational/teleport/lib/relaytunnel"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/scopes"
authorizedkeysreporter "github.com/gravitational/teleport/lib/secretsscanner/authorizedkeys"
"github.com/gravitational/teleport/lib/service/servicecfg"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/srv"
"github.com/gravitational/teleport/lib/srv/ingress"
"github.com/gravitational/teleport/lib/sshagent"
"github.com/gravitational/teleport/lib/sshutils"
reexecutils "github.com/gravitational/teleport/lib/sshutils/reexec"
"github.com/gravitational/teleport/lib/utils"
hostuser "github.com/gravitational/teleport/session/host/user"
"github.com/gravitational/teleport/session/networking"
"github.com/gravitational/teleport/session/networking/x11"
"github.com/gravitational/teleport/session/pam/pamcfg"
"github.com/gravitational/teleport/session/reexec"
"github.com/gravitational/teleport/session/reexec/reexecconstants"
)
// Server implements SSH server that uses configuration backend and
// certificate-based authentication
type Server struct {
sync.Mutex
logger *slog.Logger
namespace string
addr utils.NetAddr
hostname string
srv *sshutils.Server
getRotation services.RotationGetter
authService srv.AccessPoint
reg *srv.SessionRegistry
limiter *limiter.Limiter
inventoryHandle inventory.DownstreamHandle
// labels are static labels.
labels map[string]string
// dynamicLabels are the result of command execution.
dynamicLabels *labels.Dynamic
// cloudLabels are the labels imported from a cloud provider.
cloudLabels labels.Importer
proxyMode bool
proxyClusterGetter reversetunnelclient.ClusterGetter
proxyAccessPoint authclient.ReadProxyAccessPoint
peerAddr string
advertiseAddr *utils.NetAddr
proxyPublicAddr utils.NetAddr
publicAddrs []utils.NetAddr
// server UUID gets generated once on the first start and never changes
// usually stored in a file inside the data dir
uuid string
// this gets set to true for unit testing
isTestStub bool
// testLoginShell overrides the shell used for sessions. It is only set by
// tests to avoid running the real user's shell and polluting shell history.
testLoginShell string
// cancel cancels all operations
cancel context.CancelFunc
// ctx is broadcasting context closure
ctx context.Context
// StreamEmitter points to the auth service and emits audit events
events.StreamEmitter
// clock is a system clock
clock clockwork.Clock
// permitUserEnvironment controls if this server will read ~/.tsh/environment
// before creating a new session.
permitUserEnvironment bool
// ciphers is a list of ciphers that the server supports. If omitted,
// the defaults will be used.
ciphers []string
// kexAlgorithms is a list of key exchange (KEX) algorithms that the
// server supports. If omitted, the defaults will be used.
kexAlgorithms []string
// macAlgorithms is a list of message authentication codes (MAC) that
// the server supports. If omitted the defaults will be used.
macAlgorithms []string
// authHandlers are common authorization and authentication related handlers.
authHandlers *srv.AuthHandlers
// termHandlers are common terminal related handlers.
termHandlers *srv.TermHandlers
// pamConfig holds configuration for PAM.
pamConfig *pamcfg.PAMConfig
// dataDir is a server local data directory
dataDir string
// heartbeat sends updates about this server
// back to auth server
heartbeat srv.HeartbeatI
// useTunnel is used to inform other components that this server is
// requesting connections to it come over a reverse tunnel.
useTunnel bool
// fips means Teleport started in a FedRAMP/FIPS compliant
// configuration.
fips bool
// ebpf is the service used for enhanced session recording.
ebpf bpf.BPF
// onHeartbeat is a callback for heartbeat status.
onHeartbeat func(error)
// utmpPath is the path to the user accounting database.
utmpPath string
// wtmpPath is the path to the user accounting s.Logger.
wtmpPath string
// btmpPath is the path to the user accounting failed login log.
btmpPath string
// wtmpdbPath is the path to the wtmpdb database file.
wtmpdbPath string
// allowTCPForwarding indicates whether the ssh server is allowed to offer
// TCP port forwarding.
allowTCPForwarding bool
// x11 is the X11 forwarding configuration for the server
x11 *x11.ServerConfig
// allowFileCopying indicates whether the ssh server is allowed to handle
// remote file operations via SCP or SFTP.
allowFileCopying bool
// lockWatcher is the server's lock watcher.
lockWatcher *services.LockWatcher
// connectedProxyGetter gets the proxies teleport is connected to.
connectedProxyGetter reversetunnelclient.ConnectedProxyGetter
// relayInfoGetter gets the Relay group and Relay host IDs that this
// Teleport instance is connected to. The returned data must be owned by the
// caller (i.e. it should be a copy).
relayInfoGetter relaytunnel.GetRelayInfoFunc
// createHostUser configures whether a host should allow host user
// creation
createHostUser bool
storage services.PresenceInternal
// users is used to start the automatic user deletion loop
users srv.HostUsers
// sudoers is used to manage sudoers file provisioning
sudoers srv.HostSudoers
// tracerProvider is used to create tracers capable
// of starting spans.
tracerProvider oteltrace.TracerProvider
// router used by subsystem requests to connect to nodes
// and clusters
router *proxy.Router
// sessionController is used to restrict new sessions
// based on locks and cluster preferences
sessionController *srv.SessionController
// ingressReporter reports new and active connections.
ingressReporter *ingress.Reporter
// ingressService the service name passed to the ingress reporter.
ingressService string
// proxySigner is used to generate signed PROXYv2 header so we can securely propagate client IP
proxySigner PROXYHeaderSigner
// remoteForwardingMap holds the remote port forwarding listeners that need
// to be closed when forwarding finishes, keyed by listen addr.
remoteForwardingMap utils.SyncMap[remoteForwardingMapKey, io.Closer]
// stableUnixUsers is used to obtain fallback UIDs for host user
// provisioning from the control plane.
stableUnixUsers stableunixusersv1.StableUNIXUsersServiceClient
// enableSELinux configures whether SELinux support is enable or not.
enableSELinux bool
// scope is the scope the server is constrained to
scope string
// childLogConfig is the log config for child processes.
childLogConfig *srv.ChildLogConfig
// immutableLabels are the immutable labels assigned to the server's host certificate
immutableLabels map[string]string
// presenceMaxDuration is the max duration that a moderated session
// can continue between presence verifications.
presenceMaxDuration time.Duration
}
type remoteForwardingMapKey struct {
user string
cluster string
srcAddr string
}
func getRemoteForwardingMapKey(scx *srv.ServerContext) remoteForwardingMapKey {
return remoteForwardingMapKey{
user: scx.Identity.TeleportUser,
cluster: scx.Identity.OriginClusterName,
srcAddr: scx.SrcAddr,
}
}
// EventMetadata returns metadata about the server.
func (s *Server) EventMetadata() apievents.ServerMetadata {
serverInfo := s.GetInfo()
return apievents.ServerMetadata{
ServerVersion: teleport.Version,
ServerNamespace: serverInfo.GetNamespace(),
ServerID: serverInfo.GetName(),
ServerAddr: serverInfo.GetAddr(),
ServerLabels: serverInfo.GetAllLabels(),
ServerHostname: serverInfo.GetHostname(),
ServerSubKind: serverInfo.GetSubKind(),
}
}
// GetClock returns server clock implementation
func (s *Server) GetClock() clockwork.Clock {
return s.clock
}
// GetDataDir returns server data dir
func (s *Server) GetDataDir() string {
return s.dataDir
}
func (s *Server) GetNamespace() string {
return s.namespace
}
func (s *Server) GetAccessPoint() srv.AccessPoint {
return s.authService
}
// GetUserAccountingPaths returns the optional override of the utmp, wtmp, and btmp paths.
func (s *Server) GetUserAccountingPaths() (utmp string, wtmp string, btmp string, wtmpdb string) {
return s.utmpPath, s.wtmpPath, s.btmpPath, s.wtmpdbPath
}
// GetPAM returns the PAM configuration for this server.
func (s *Server) GetPAM() *pamcfg.PAMConfig {
return s.pamConfig
}
// UseTunnel used to determine if this node has connected to this cluster
// using reverse tunnel.
func (s *Server) UseTunnel() bool {
return s.useTunnel
}
// GetBPF returns the BPF service used by enhanced session recording.
func (s *Server) GetBPF() bpf.BPF {
return s.ebpf
}
// GetLockWatcher gets the server's lock watcher.
func (s *Server) GetLockWatcher() *services.LockWatcher {
return s.lockWatcher
}
// GetCreateHostUser determines whether users should be created on the
// host automatically
func (s *Server) GetCreateHostUser() bool {
// we shouldn't allow creating host users on a proxy server
return !s.proxyMode && s.createHostUser
}
// GetHostUsers returns the HostUsers instance being used to manage
// host user provisioning
func (s *Server) GetHostUsers() srv.HostUsers {
return s.users
}
// GetHostSudoers returns the HostSudoers instance being used to manage
// sudoers file provisioning
func (s *Server) GetHostSudoers() srv.HostSudoers {
// we shouldn't allow modifying sudoers on a proxy server
if s.proxyMode {
return nil
}
if s.sudoers == nil {
return &srv.HostSudoersNotImplemented{}
}
return s.sudoers
}
// GetSELinuxEnabled returns whether the node should enable SELinux
// support or not.
func (s *Server) GetSELinuxEnabled() bool {
return s.enableSELinux
}
// GetProxyMode returns whether the server is started in SSH proxying mode.
func (s *Server) GetProxyMode() bool {
return s.proxyMode
}
// ChildLogConfig returns the child log config.
func (s *Server) ChildLogConfig() srv.ChildLogConfig {
if s.childLogConfig != nil {
return *s.childLogConfig
}
// return a noop log configuration
return srv.ChildLogConfig{
ExecLogConfig: reexec.ExecLogConfig{},
Writer: io.Discard,
}
}
// GetPresenceMaxDuration returns the max duration that a moderated session
// can continue between presence verifications.
func (s *Server) GetPresenceMaxDuration() time.Duration {
return s.presenceMaxDuration
}
// ServerOption is a functional option passed to the server
type ServerOption func(s *Server) error
func (s *Server) close() {
s.cancel()
s.reg.Close()
if s.heartbeat != nil {
if err := s.heartbeat.Close(); err != nil {
s.logger.WarnContext(s.ctx, "Failed to close heartbeat", "error", err)
}
}
if s.dynamicLabels != nil {
s.dynamicLabels.Close()
}
if s.users != nil {
s.users.Shutdown()
}
}
// Close closes listening socket and stops accepting connections
func (s *Server) Close() error {
s.close()
// Close the server first so we don't accept any new forwarding connections
// after we've closed them all.
errors := []error{s.srv.Close()}
s.remoteForwardingMap.Range(func(_ remoteForwardingMapKey, closer io.Closer) bool {
if closer != nil {
if err := closer.Close(); err != nil {
errors = append(errors, err)
}
}
return true
})
return trace.NewAggregate(errors...)
}
// Shutdown performs graceful shutdown
func (s *Server) Shutdown(ctx context.Context) error {
// Stop heart beating immediately to prevent active connections
// from making the server appear alive and well.
if s.heartbeat != nil {
if err := s.heartbeat.Close(); err != nil {
s.logger.WarnContext(ctx, "Failed to close heartbeat", "error", err)
}
}
// wait until connections drain off
err := s.srv.Shutdown(ctx)
return trace.NewAggregate(err, s.Close())
}
// Start starts server
func (s *Server) Start() error {
// Only call srv.Start() which listens on a socket if the server did not
// request connections to it arrive over a reverse tunnel.
if !s.useTunnel {
if err := s.srv.Start(); err != nil {
return trace.Wrap(err)
}
}
// Heartbeat should start only after s.srv.Start.
// If the server is configured to listen on port 0 (such as in tests),
// it'll only populate its actual listening address during s.srv.Start.
// Heartbeat uses this address to announce. Avoid announcing an empty
// address on first heartbeat.
s.startPeriodicOperations()
return nil
}
// Serve servers service on started listener
func (s *Server) Serve(l net.Listener) error {
// Set the listener before starting heartbeats so the first node heartbeat
// does not advertise an empty address.
if err := s.srv.SetListener(l); err != nil {
return trace.Wrap(err)
}
s.startPeriodicOperations()
return trace.Wrap(s.srv.Serve())
}
func (s *Server) startPeriodicOperations() {
// If the server has dynamic labels defined, start a loop that will
// asynchronously keep them updated.
if s.dynamicLabels != nil {
go s.dynamicLabels.Start()
}
// If the server allows host user provisioning, this will start an
// automatic cleanup process for any temporary leftover users.
if s.GetCreateHostUser() && s.users != nil {
go s.users.UserCleanup()
}
if s.cloudLabels != nil {
s.cloudLabels.Start(s.Context())
}
if s.heartbeat != nil {
go s.heartbeat.Run()
}
}
// Wait waits until server stops
func (s *Server) Wait() {
s.srv.Wait(context.TODO())
}
// HandleConnection is called after a connection has been accepted and starts
// to perform the SSH handshake immediately.
func (s *Server) HandleConnection(conn net.Conn) {
s.srv.HandleConnection(conn)
}
// SetUserAccountingPaths is a functional server option to override the user accounting database and log path.
func SetUserAccountingPaths(utmpPath, wtmpPath, btmpPath, wtmpdbPath string) ServerOption {
return func(s *Server) error {
s.utmpPath = utmpPath
s.wtmpPath = wtmpPath
s.btmpPath = btmpPath
s.wtmpdbPath = wtmpdbPath
return nil
}
}
// SetClock is a functional server option to override the internal
// clock
func SetClock(clock clockwork.Clock) ServerOption {
return func(s *Server) error {
s.clock = clock
return nil
}
}
// SetRotationGetter sets rotation state getter
func SetRotationGetter(getter services.RotationGetter) ServerOption {
return func(s *Server) error {
s.getRotation = getter
return nil
}
}
// SetProxyMode starts this server in SSH proxying mode
func SetProxyMode(peerAddr string, clusterGetter reversetunnelclient.ClusterGetter, ap authclient.ReadProxyAccessPoint, router *proxy.Router) ServerOption {
return func(s *Server) error {
// always set proxy mode to true,
// because in some tests reverse tunnel is disabled,
// but proxy is still used without it.
s.proxyMode = true
s.proxyClusterGetter = clusterGetter
s.proxyAccessPoint = ap
s.peerAddr = peerAddr
s.router = router
return nil
}
}
// SetIngressReporter sets the reporter for reporting new and active connections.
func SetIngressReporter(service string, r *ingress.Reporter) ServerOption {
return func(s *Server) error {
s.ingressReporter = r
s.ingressService = service
return nil
}
}
// SetLabels sets dynamic and static labels that server will report to the
// auth servers.
func SetLabels(staticLabels map[string]string, cmdLabels services.CommandLabels, cloudLabels labels.Importer) ServerOption {
return func(s *Server) error {
var err error
// clone and validate labels and cmdLabels. in theory,
// only cmdLabels should experience concurrent writes,
// but this operation is only run once on startup
// so a little defensive cloning is harmless.
labelsClone := make(map[string]string, len(staticLabels))
for name, label := range staticLabels {
if !types.IsValidLabelKey(name) {
return trace.BadParameter("invalid label key: %q", name)
}
labelsClone[name] = label
}
s.labels = labelsClone
if len(cmdLabels) > 0 {
// Create dynamic labels.
s.dynamicLabels, err = labels.NewDynamic(s.ctx, &labels.DynamicConfig{
Labels: cmdLabels,
})
if err != nil {
return trace.Wrap(err)
}
}
s.cloudLabels = cloudLabels
return nil
}
}
// SetLimiter sets rate and connection limiter for this server
func SetLimiter(limiter *limiter.Limiter) ServerOption {
return func(s *Server) error {
s.limiter = limiter
return nil
}
}
// SetEmitter assigns an audit event emitter for this server
func SetEmitter(emitter events.StreamEmitter) ServerOption {
return func(s *Server) error {
s.StreamEmitter = emitter
return nil
}
}
// SetUUID sets server unique ID
func SetUUID(uuid string) ServerOption {
return func(s *Server) error {
s.uuid = uuid
return nil
}
}
func SetNamespace(namespace string) ServerOption {
return func(s *Server) error {
s.namespace = namespace
return nil
}
}
// SetPermitUserEnvironment allows you to set the value of permitUserEnvironment.
func SetPermitUserEnvironment(permitUserEnvironment bool) ServerOption {
return func(s *Server) error {
s.permitUserEnvironment = permitUserEnvironment
return nil
}
}
func SetCiphers(ciphers []string) ServerOption {
return func(s *Server) error {
s.ciphers = ciphers
return nil
}
}
func SetKEXAlgorithms(kexAlgorithms []string) ServerOption {
return func(s *Server) error {
s.kexAlgorithms = kexAlgorithms
return nil
}
}
func SetMACAlgorithms(macAlgorithms []string) ServerOption {
return func(s *Server) error {
s.macAlgorithms = macAlgorithms
return nil
}
}
func SetPAMConfig(pamConfig *pamcfg.PAMConfig) ServerOption {
return func(s *Server) error {
s.pamConfig = pamConfig
return nil
}
}
func SetUseTunnel(useTunnel bool) ServerOption {
return func(s *Server) error {
s.useTunnel = useTunnel
return nil
}
}
func SetFIPS(fips bool) ServerOption {
return func(s *Server) error {
s.fips = fips
return nil
}
}
func SetBPF(ebpf bpf.BPF) ServerOption {
return func(s *Server) error {
s.ebpf = ebpf
return nil
}
}
func SetOnHeartbeat(fn func(error)) ServerOption {
return func(s *Server) error {
s.onHeartbeat = fn
return nil
}
}
// SetCreateHostUser configures host user creation on a server
func SetCreateHostUser(createUser bool) ServerOption {
return func(s *Server) error {
s.createHostUser = createUser && runtime.GOOS == constants.LinuxOS
return nil
}
}
// SetStoragePresenceService configures host user creation on a server
func SetStoragePresenceService(service services.PresenceInternal) ServerOption {
return func(s *Server) error {
s.storage = service
return nil
}
}
// SetAllowTCPForwarding sets the TCP port forwarding mode that this server is
// allowed to offer. The default value is SSHPortForwardingModeAll, i.e. port
// forwarding is allowed.
func SetAllowTCPForwarding(allow bool) ServerOption {
return func(s *Server) error {
s.allowTCPForwarding = allow
return nil
}
}
// SetLockWatcher sets the server's lock watcher.
func SetLockWatcher(lockWatcher *services.LockWatcher) ServerOption {
return func(s *Server) error {
s.lockWatcher = lockWatcher
return nil
}
}
// SetX11ForwardingConfig sets the server's X11 forwarding configuration
func SetX11ForwardingConfig(xc *x11.ServerConfig) ServerOption {
return func(s *Server) error {
s.x11 = xc
return nil
}
}
// SetAllowFileCopying sets whether the server is allowed to handle
// SCP/SFTP requests.
func SetAllowFileCopying(allow bool) ServerOption {
return func(s *Server) error {
s.allowFileCopying = allow
return nil
}
}
// SetConnectedProxyGetter sets the ConnectedProxyGetter.
func SetConnectedProxyGetter(getter reversetunnelclient.ConnectedProxyGetter) ServerOption {
return func(s *Server) error {
s.connectedProxyGetter = getter
return nil
}
}
// SetRelayInfoGetter sets the function used to get the relay tunnel client info
// to fill in the server heartbeat.
func SetRelayInfoGetter(getter relaytunnel.GetRelayInfoFunc) ServerOption {
return func(s *Server) error {
s.relayInfoGetter = getter
return nil
}
}
// SetInventoryControlHandle sets the server's downstream inventory control
// handle.
func SetInventoryControlHandle(handle inventory.DownstreamHandle) ServerOption {
return func(s *Server) error {
s.inventoryHandle = handle
return nil
}
}
// SetTracerProvider sets the tracer provider.
func SetTracerProvider(provider oteltrace.TracerProvider) ServerOption {
return func(s *Server) error {
s.tracerProvider = provider
return nil
}
}
// SetSessionController sets the session controller.
func SetSessionController(controller *srv.SessionController) ServerOption {
return func(s *Server) error {
s.sessionController = controller
return nil
}
}
// SetPROXYSigner sets the PROXY headers signer
func SetPROXYSigner(proxySigner PROXYHeaderSigner) ServerOption {
return func(s *Server) error {
s.proxySigner = proxySigner
return nil
}
}
// SetStableUNIXUsers sets the client for the stable UNIX users service, used as
// a fallback to get UIDs for host user creation.
func SetStableUNIXUsers(stableUNIXUsers stableunixusersv1.StableUNIXUsersServiceClient) ServerOption {
return func(s *Server) error {
s.stableUnixUsers = stableUNIXUsers
return nil
}
}
// GetSELinuxEnabled returns whether the node should enable SELinux
// support or not.
func SetSELinuxEnabled(enabled bool) ServerOption {
return func(s *Server) error {
s.enableSELinux = enabled
return nil
}
}
// SetPublicAddrs sets the server's public addresses
func SetPublicAddrs(addrs []utils.NetAddr) ServerOption {
return func(s *Server) error {
s.publicAddrs = addrs
return nil
}
}
// SetScope sets the server's scope.
func SetScope(scope string) ServerOption {
return func(s *Server) error {
if scope == "" {
s.scope = ""
return nil
}
if err := scopes.WeakValidate(scope); err != nil {
return trace.Wrap(err)
}
s.scope = scope
return nil
}
}
// SetImmutableLabels sets the server's immutable labels.
func SetImmutableLabels(labels map[string]string) ServerOption {
return func(s *Server) error {
s.immutableLabels = labels
return nil
}
}
// SetChildLogConfig sets the config that will be used to handle logs
// from child processes.
func SetChildLogConfig(cfg *servicecfg.Config) ServerOption {
return func(s *Server) error {
s.childLogConfig = &srv.ChildLogConfig{
ExecLogConfig: reexec.ExecLogConfig{
Level: cfg.LoggerLevel.Level(),
Format: strings.ToLower(cfg.LogConfig.Format),
ExtraFields: cfg.LogConfig.ExtraFields,
EnableColors: cfg.LogConfig.EnableColors,
Padding: cfg.LogConfig.Padding,
},
Writer: cfg.LogWriter,
}
return nil
}
}
// SetPresenceMaxDuration sets the max duration that a moderated session
// can continue between presence verifications.
func SetPresenceMaxDuration(maxDuration time.Duration) ServerOption {
return func(s *Server) error {
s.presenceMaxDuration = maxDuration
return nil
}
}
// SetTestLoginShell overrides the shell used for sessions instead of resolving
// the login shell of the OS user. This is intended for tests only to avoid
// running the real user's shell and polluting shell history.
func SetTestLoginShell(shell string) ServerOption {
return func(s *Server) error {
s.testLoginShell = shell
return nil
}
}
// New returns an unstarted server
func New(
ctx context.Context,
addr utils.NetAddr,
hostname string,
getHostSigners sshutils.GetHostSignersFunc,
authService srv.AccessPoint,
dataDir string,
advertiseAddr string,
proxyPublicAddr utils.NetAddr,
auth authclient.ClientI,
options ...ServerOption,
) (*Server, error) {
ctx, cancel := context.WithCancel(ctx)
s := &Server{
addr: addr,
authService: authService,
hostname: hostname,
proxyPublicAddr: proxyPublicAddr,
cancel: cancel,
ctx: ctx,
clock: clockwork.NewRealClock(),
dataDir: dataDir,
allowTCPForwarding: true,
}
var err error
s.limiter, err = limiter.NewLimiter(limiter.Config{})
if err != nil {
return nil, trace.Wrap(err)
}
if advertiseAddr != "" {
s.advertiseAddr, err = utils.ParseAddr(advertiseAddr)
if err != nil {
return nil, trace.Wrap(err)
}
}
for _, o := range options {
if err := o(s); err != nil {
return nil, trace.Wrap(err)
}
}
if s.uuid == "" {
return nil, trace.BadParameter("server UUID must be set using SetUUID")
}
// TODO(klizhentas): replace function arguments with struct
if s.StreamEmitter == nil {
return nil, trace.BadParameter("setup valid Emitter parameter using SetEmitter")
}
if s.namespace == "" {
return nil, trace.BadParameter("setup valid namespace parameter using SetNamespace")
}
if s.lockWatcher == nil {
return nil, trace.BadParameter("setup valid LockWatcher parameter using SetLockWatcher")
}
if s.sessionController == nil {
return nil, trace.BadParameter("setup valid SessionControl parameter using SetSessionControl")
}
if s.connectedProxyGetter == nil {
return nil, trace.BadParameter("setup valid ConnectedProxyGetter parameter using SetConnectedProxyGetter")
}
if s.tracerProvider == nil {
s.tracerProvider = tracing.DefaultProvider()
}
if s.presenceMaxDuration == 0 {
s.presenceMaxDuration = client.DefaultPresenceMaxDuration
}
var component string
if s.proxyMode {
component = teleport.ComponentProxy
} else {
component = teleport.ComponentNode
}
s.logger = slog.With(teleport.ComponentKey, component)
if s.GetCreateHostUser() {
s.users = srv.NewHostUsers(ctx, s.storage, s.ID())
}
s.sudoers = srv.NewHostSudoers(s.ID())
s.reg, err = srv.NewSessionRegistry(srv.SessionRegistryConfig{
Srv: s,
SessionTrackerService: auth,
})
if err != nil {
return nil, trace.Wrap(err)
}
// add in common auth handlers
authHandlerConfig := srv.AuthHandlerConfig{
Server: s,
Component: component,
AccessPoint: s.authService,
FIPS: s.fips,
Emitter: s.StreamEmitter,
Clock: s.clock,
ValidatedMFAChallengeVerifier: auth.MFAServiceClientV2(),
}
s.authHandlers, err = srv.NewAuthHandlers(&authHandlerConfig)
if err != nil {
return nil, trace.Wrap(err)
}
// common term handlers
s.termHandlers = &srv.TermHandlers{
SessionRegistry: s.reg,
}
clusterName, err := s.GetAccessPoint().GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
server, err := sshutils.NewServer(
component,
addr, s,
getHostSigners,
sshutils.AuthMethods{
PublicKey: s.authHandlers.PublicKeyCallback,
VerifiedPublicKey: s.authHandlers.VerifiedPublicKeyCallback,
},
sshutils.SetLimiter(s.limiter),
sshutils.SetRequestHandler(s),
sshutils.SetNewConnHandler(s),
sshutils.SetCiphers(s.ciphers),
sshutils.SetKEXAlgorithms(s.kexAlgorithms),
sshutils.SetMACAlgorithms(s.macAlgorithms),
sshutils.SetFIPS(s.fips),
sshutils.SetClock(s.clock),
sshutils.SetIngressReporter(s.ingressService, s.ingressReporter),
sshutils.SetClusterName(clusterName.GetClusterName()),
)
if err != nil {
return nil, trace.Wrap(err)
}
s.srv = server
if !s.proxyMode {
if err := s.startAuthorizedKeysManager(ctx, auth); err != nil {
s.logger.InfoContext(ctx, "Failed to start authorized keys manager", "error", err)
}
}
var heartbeat srv.HeartbeatI
if !s.proxyMode {
if s.inventoryHandle == nil {
return nil, trace.BadParameter("inventoryHandle must not be nil")
}
s.logger.DebugContext(ctx, "starting control-stream based heartbeat")
heartbeat, err = srv.NewSSHServerHeartbeat(srv.HeartbeatV2Config[*types.ServerV2]{
InventoryHandle: s.inventoryHandle,
GetResource: s.getServerInfo,
OnHeartbeat: s.onHeartbeat,
})
} else {
s.logger.DebugContext(ctx, "starting legacy heartbeat")
heartbeat, err = srv.NewHeartbeat(srv.HeartbeatConfig{
Mode: srv.HeartbeatModeProxy,
Context: ctx,
Component: component,
Announcer: s.authService,
GetServerInfo: s.getServerResource,
KeepAlivePeriod: apidefaults.ServerKeepAliveTTL(),
AnnouncePeriod: apidefaults.ProxyAnnounceTTL()/2 + utils.RandomDuration(apidefaults.ProxyAnnounceTTL()/10),
ServerTTL: apidefaults.ProxyAnnounceTTL(),
CheckPeriod: defaults.HeartbeatCheckPeriod,
Clock: s.clock,
OnHeartbeat: s.onHeartbeat,
})
}
if err != nil {
s.srv.Close()
return nil, trace.Wrap(err)
}
s.heartbeat = heartbeat
return s, nil
}
func (s *Server) getNamespace() string {
return types.ProcessNamespace(s.namespace)
}
func (s *Server) clusterGetterWithAccessChecker(ctx *srv.ServerContext) (reversetunnelclient.ClusterGetter, error) {
clusterName, err := s.GetAccessPoint().GetClusterName(s.ctx)
if err != nil {
return nil, trace.Wrap(err)
}
return reversetunnelclient.NewClusterGetterWithRoles(s.proxyClusterGetter, clusterName.GetClusterName(), ctx.Identity.UnstableClusterAccessChecker, s.proxyAccessPoint), nil
}
// startAuthorizedKeysManager starts the authorized keys manager.
func (s *Server) startAuthorizedKeysManager(ctx context.Context, auth authclient.ClientI) error {
authorizedKeysWatcher, err := authorizedkeysreporter.NewWatcher(
ctx,
authorizedkeysreporter.WatcherConfig{
Client: auth,
Logger: slog.Default(),
HostID: s.uuid,
Clock: s.clock,
},
)
if errors.Is(err, hostuser.ErrUnsupportedPlatform) {
return nil
} else if err != nil {
return trace.Wrap(err)
}
go func() {
if err := authorizedKeysWatcher.Run(ctx); err != nil {
s.logger.WarnContext(ctx, "Failed to start authorized keys watcher", "error", err)
}
}()
return nil
}
// Context returns server shutdown context
func (s *Server) Context() context.Context {
return s.ctx
}
func (s *Server) Component() string {
if s.proxyMode {
return teleport.ComponentProxy
}
return teleport.ComponentNode
}
// Addr returns server address
func (s *Server) Addr() string {
return s.srv.Addr()
}
// ActiveConnections returns the number of connections that are
// being served.
func (s *Server) ActiveConnections() int32 {
return s.srv.ActiveConnections()
}
// ID returns server ID
func (s *Server) ID() string {
return s.uuid
}
// HostUUID is the ID of the server. This value is the same as ID, it is
// different from the forwarding server.
func (s *Server) HostUUID() string {
return s.uuid
}
// PermitUserEnvironment returns if ~/.tsh/environment will be read before a
// session is created by this server.
func (s *Server) PermitUserEnvironment() bool {
return s.permitUserEnvironment
}
func (s *Server) setAdvertiseAddr(addr *utils.NetAddr) {
s.Lock()
defer s.Unlock()
s.advertiseAddr = addr
}
func (s *Server) getAdvertiseAddr() *utils.NetAddr {
s.Lock()
defer s.Unlock()
return s.advertiseAddr
}
// AdvertiseAddr returns an address this server should be publicly accessible
// as, in "ip:host" form
func (s *Server) AdvertiseAddr() string {
// set if we have explicit --advertise-ip option
advertiseAddr := s.getAdvertiseAddr()
listenAddr := s.Addr()
if advertiseAddr == nil {
return listenAddr
}
_, port, _ := net.SplitHostPort(listenAddr)
ahost, aport, err := utils.ParseAdvertiseAddr(advertiseAddr.String())
if err != nil {
s.logger.WarnContext(s.ctx, "Failed to parse advertise address, using default value", "advertise_addr", advertiseAddr, "error", err, "default_addr", listenAddr)
return listenAddr
}
if aport == "" {
aport = port
}
return net.JoinHostPort(ahost, aport)
}
func (s *Server) getRole() types.SystemRole {
if s.proxyMode {
return types.RoleProxy
}
return types.RoleNode
}
// getStaticLabels gets the labels that the server should present as static,
// which includes EC2 labels if available.
func (s *Server) getStaticLabels() map[string]string {
labels := make(map[string]string, len(s.labels))
if s.cloudLabels != nil {
maps.Copy(labels, s.cloudLabels.Get())
}
// Let labels sent over ics override labels from instance metadata.
if s.inventoryHandle != nil {
maps.Copy(labels, s.inventoryHandle.GetUpstreamLabels(proto.LabelUpdateKind_SSHServerCloudLabels))
}
// Let static labels override any other labels.
maps.Copy(labels, s.labels)
return labels
}
// getDynamicLabels returns all dynamic labels. If no dynamic labels are
// defined, return an empty set.
func (s *Server) getDynamicLabels() map[string]types.CommandLabelV2 {
if s.dynamicLabels == nil {
return make(map[string]types.CommandLabelV2)
}
return types.LabelsToV2(s.dynamicLabels.Get())
}
// GetInfo returns a services.Server that represents this server.
func (s *Server) GetInfo() types.Server {
return s.getBasicInfo()
}
func (s *Server) getBasicInfo() *types.ServerV2 {
// Only set the address for non-tunnel nodes.
var addr string
if !s.useTunnel {
addr = s.AdvertiseAddr()
}
var relayGroup string
var relayIDs []string
if s.relayInfoGetter != nil {
// relayInfoGetter returns a copy of the slice, so we can move it in the
// protobuf message
relayGroup, relayIDs = s.relayInfoGetter()
}
kind := types.KindNode
if s.proxyMode {
kind = types.KindProxy
}
srv := &types.ServerV2{
Kind: kind,
Version: types.V2,
Scope: s.scope,
Metadata: types.Metadata{
Name: s.ID(),
Namespace: s.getNamespace(),
Labels: s.getStaticLabels(),
},
Spec: types.ServerSpecV2{
CmdLabels: s.getDynamicLabels(),
Addr: addr,
Hostname: s.hostname,
UseTunnel: s.useTunnel,
Version: teleport.Version,
ProxyIDs: s.connectedProxyGetter.GetProxyIDs(),
RelayGroup: relayGroup,
RelayIds: relayIDs,
ImmutableLabels: s.immutableLabels,
},
}
srv.SetPublicAddrs(utils.NetAddrsToStrings(s.publicAddrs))
srv.SetComponentFeatures(componentfeatures.ForSSHServer())
return srv
}
func (s *Server) getServerInfo(ctx context.Context) (*types.ServerV2, error) {
server := s.getBasicInfo()
if s.getRotation != nil {
rotation, err := s.getRotation(s.getRole())
if err != nil {
if !trace.IsNotFound(err) {
s.logger.WarnContext(ctx, "Failed to get rotation state", "error", err)
}
} else {
server.SetRotation(*rotation)
}
}
server.SetExpiry(s.clock.Now().UTC().Add(apidefaults.ServerAnnounceTTL))
server.SetPeerAddr(s.peerAddr)
return server, nil
}
func (s *Server) getServerResource() (types.Resource, error) {
return s.getServerInfo(s.ctx)
}
// dialTCPIP dials the given tcpip address through the network process.
func (s *Server) dialTCPIP(ctx context.Context, scx *srv.ServerContext, addr string) (net.Conn, error) {
proc, err := s.getNetworkingProcess(ctx, scx)
if err != nil {
return nil, trace.Wrap(err)
}
conn, err := proc.Dial(ctx, "tcp", addr)
if err != nil {
return nil, trace.Wrap(err)
}
return conn, nil
}
// listenTCPIP creates a new listener in the networking process.
func (s *Server) listenTCPIP(ctx context.Context, scx *srv.ServerContext, addr string) (net.Listener, error) {
proc, err := s.getNetworkingProcess(ctx, scx)
if err != nil {
return nil, trace.Wrap(err)
}
listener, err := proc.Listen(ctx, "tcp", addr)
if err != nil {
return nil, trace.Wrap(err)
}
return listener, nil
}
// getNetworkingProcess sets up a connection-level subprocess that handles
// networking requests. Subsequent calls from the same connection context
// reuse the same networking process.
func (s *Server) getNetworkingProcess(ctx context.Context, scx *srv.ServerContext) (*networking.Process, error) {
if proc, ok := scx.Parent().GetNetworkingProcess(); ok {
return proc, nil
}
proc, err := s.startNetworkingProcess(ctx, scx)
if err != nil {
return nil, trace.Wrap(err)
}
// Try to register with the parent context.
if otherProc, ok := scx.Parent().SetNetworkingProcess(proc); !ok {
// Another networking process was concurrently created. this isn't actually a problem, multiple networking
// processes being registered is harmless, but it does result in slightly higher resource utilization, so
// it's preferable to use the existing networking process and close ours in the background.
go proc.Close()
return otherProc, nil
}
scx.Parent().AddCloser(proc)
return proc, nil
}
// startNetworkingProcess launches a new networking process. It should be closed once
// the server connection is closed.
func (s *Server) startNetworkingProcess(ctx context.Context, scx *srv.ServerContext) (*networking.Process, error) {
// Create context for the networking process.
nsctx, err := srv.NewServerContext(ctx, scx.ConnectionContext, s, scx.Identity, nil)
if err != nil {
return nil, trace.Wrap(err)
}
nsctx.SessionRecordingConfig.SetMode(types.RecordOff)
nsctx.ExecType = reexecconstants.NetworkingSubCommand
scx.Parent().AddCloser(nsctx)
// Create command to re-exec Teleport which will handle networking requests. The
// reason it's not done directly is because the PAM stack needs to be called
// from the child process.
cmd, err := nsctx.ConfigureCommand(nil)
if err != nil {
return nil, trace.Wrap(err)
}
proc, childErr, err := networking.NewProcess(ctx, cmd.Cmd)
if err != nil {
if childErr == "" {
return nil, trace.Wrap(err)
}
// If the networking process failed with an error message from stderr, prefer
// that over the other error.
childErr = reexecutils.ChildErrorWithContext(childErr, &reexecutils.ErrorContext{
DecisionContext: scx.Identity.AccessPermit.GetDecisionContext(),
Login: scx.Identity.Login,
})
return nil, errors.New(strings.TrimRight(childErr, "\n"))
}
return proc, nil
}
// HandleRequest processes global out-of-band requests. Global out-of-band
// requests are processed in order (this way the originator knows which
// request we are responding to). If Teleport does not support the request
// type or an error occurs while processing that request Teleport will reply
// req.Reply(false, nil).
//
// For more details: https://tools.ietf.org/html/rfc4254.html#page-4
func (s *Server) HandleRequest(ctx context.Context, ccx *sshutils.ConnectionContext, r *ssh.Request) {
switch r.Type {
case teleport.KeepAliveReqType:
s.handleKeepAlive(r)
case teleport.ClusterDetailsReqType:
s.handleClusterDetails(ctx, r)
case teleport.VersionRequest:
s.handleVersionRequest(ctx, r)
case teleport.TerminalSizeRequest:
if err := s.termHandlers.HandleTerminalSize(r); err != nil {
s.logger.WarnContext(ctx, "failed to handle terminal size request", "error", err)
if r.WantReply {
if err := r.Reply(false, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to terminal size request", "error", err)
}
}
}
case teleport.TCPIPForwardRequest:
if err := s.handleTCPIPForwardRequest(ctx, ccx, r); err != nil {
s.logger.WarnContext(ctx, "failed to handle tcpip forward request", "error", err)
if err := r.Reply(false, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to tcpip forward request", "error", err)
}
}
case teleport.CancelTCPIPForwardRequest:
if err := s.handleCancelTCPIPForwardRequest(ctx, ccx, r); err != nil {
s.logger.WarnContext(ctx, "failed to handle cancel tcpip forward request", "error", err)
if err := r.Reply(false, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to tcpip forward request", "error", err)
}
}
case teleport.SessionIDQueryRequest:
// TODO(Joerger): DELETE IN v20.0.0
// All v17+ servers set the session ID. v19+ clients stop checking.
// Reply true to session ID query requests, we will set new
// session IDs for new sessions during the shel/exec channel
// request.
if err := r.Reply(true, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to session ID query request", "error", err)
}
return
case teleport.SessionIDQueryRequestV2:
// TODO(Joerger): DELETE IN v21.0.0
// clients should stop checking in v21, and servers should stop responding to the query in v22.
// Reply true to session ID query requests, we will set new
// session IDs for new sessions directly after accepting the
// session channel request.
if err := r.Reply(true, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to session ID query request", "error", err)
}
return
default:
if err := r.Reply(false, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to ssh request", "request_type", r.Type, "error", err)
}
s.logger.DebugContext(ctx, "Discarding global request", "request_type", r.Type)
}
}
// HandleNewConn is called by sshutils.Server once for each new incoming connection,
// prior to handling any channels or requests.
func (s *Server) HandleNewConn(ctx context.Context, ccx *sshutils.ConnectionContext) (context.Context, error) {
identityContext, err := s.authHandlers.CreateIdentityContext(ccx.ServerConn)
if err != nil {
return ctx, trace.Wrap(err)
}
// Apply session control restrictions.
ctx, err = s.sessionController.AcquireSessionContext(ctx, identityContext, ccx.ServerConn.LocalAddr().String(), ccx.ServerConn.RemoteAddr().String(), ccx)
if err != nil {
return ctx, trace.Wrap(err)
}
// Create host user.
created, userCloser, err := s.termHandlers.SessionRegistry.UpsertHostUser(identityContext, s.obtainFallbackUID)
if err != nil {
s.logger.WarnContext(ctx, "error while creating host users", "error", err)
}
// Indicate that the user was created by Teleport.
ccx.UserCreatedByTeleport = created
if userCloser != nil {
ccx.AddCloser(userCloser)
}
sudoersCloser, err := s.termHandlers.SessionRegistry.WriteSudoersFile(identityContext)
if err != nil {
s.logger.WarnContext(ctx, "error while writing sudoers", "error", err)
}
if sudoersCloser != nil {
ccx.AddCloser(sudoersCloser)
}
return ctx, nil
}
// obtainFallbackUID checks if the cluster is configured for stable
// autoprovisioned UNIX user UIDs and, if so, obtains and returns the UID for
// the given username. If the cluster is not configured for stable UIDs, it
// returns (_, false, nil).
func (s *Server) obtainFallbackUID(ctx context.Context, username string) (uid int32, ok bool, _ error) {
authPref, err := s.authService.GetAuthPreference(ctx)
if err != nil {
return 0, false, trace.Wrap(err)
}
cfg := authPref.GetStableUNIXUserConfig()
if cfg == nil || !cfg.Enabled {
return 0, false, nil
}
resp, err := s.stableUnixUsers.ObtainUIDForUsername(ctx, stableunixusersv1.ObtainUIDForUsernameRequest_builder{
Username: username,
}.Build())
if err != nil {
return 0, false, trace.Wrap(err)
}
uid = resp.GetUid()
// see https://github.com/systemd/systemd/blob/cc7300fc5868f6d47f3f47076100b574bf54e58d/docs/UIDS-GIDS.md
const firstUserUID = 1000
if uid < firstUserUID {
return 0, false, trace.BadParameter("received a negative or system UID as the new UID from the control plane (%v)", uid)
}
return uid, true, nil
}
// HandleNewChan is called when new channel is opened
func (s *Server) HandleNewChan(ctx context.Context, ccx *sshutils.ConnectionContext, nch ssh.NewChannel) {
identityContext, err := s.authHandlers.CreateIdentityContext(ccx.ServerConn)
if err != nil {
s.rejectChannel(ctx, nch, ssh.Prohibited, fmt.Sprintf("Unable to create identity from connection: %v", err))
return
}
channelType := nch.ChannelType()
if s.proxyMode {
switch channelType {
// Channels of type "direct-tcpip", for proxies, it's equivalent
// of teleport proxy: subsystem
case teleport.ChanDirectTCPIP:
req, err := sshutils.ParseDirectTCPIPReq(nch.ExtraData())
if err != nil {
s.logger.ErrorContext(ctx, "Failed to parse request data", "data", string(nch.ExtraData()), "error", err)
s.rejectChannel(ctx, nch, ssh.UnknownChannelType, "failed to parse direct-tcpip request")
return
}
ch, reqC, err := nch.Accept()
if err != nil {
s.logger.WarnContext(ctx, "Unable to accept channel", "error", err)
s.rejectChannel(ctx, nch, ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
return
}
go ssh.DiscardRequests(reqC)
go s.handleProxyJump(ctx, ccx, identityContext, ch, *req)
return
// Channels of type "session" handle requests that are involved in running
// commands on a server. In the case of proxy mode subsystem and agent
// forwarding requests occur over the "session" channel.
case teleport.ChanSession:
ch, requests, err := nch.Accept()
if err != nil {
s.logger.WarnContext(ctx, "Unable to accept channel", "error", err)
s.rejectChannel(ctx, nch, ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
return
}
go s.handleSessionRequests(ctx, ccx, identityContext, nil, ch, requests)
return
default:
s.rejectChannel(ctx, nch, ssh.UnknownChannelType, fmt.Sprintf("unknown channel type: %v", channelType))
return
}
}
switch channelType {
// Channels of type "session" handle requests that are involved in running
// commands on a server, subsystem requests, and agent forwarding.
case teleport.ChanSession:
var decr func()
if max := identityContext.AccessPermit.GetMaxSessions(); max != 0 {
d, ok := ccx.IncrSessions(max)
if !ok {
// user has exceeded their max concurrent ssh sessions.
if err := s.EmitAuditEvent(s.ctx, &apievents.SessionReject{
Metadata: apievents.Metadata{
Type: events.SessionRejectedEvent,
Code: events.SessionRejectedCode,
},
UserMetadata: identityContext.GetUserMetadata(),
ConnectionMetadata: apievents.ConnectionMetadata{
Protocol: events.EventProtocolSSH,
LocalAddr: ccx.ServerConn.LocalAddr().String(),
RemoteAddr: ccx.ServerConn.RemoteAddr().String(),
},
ServerMetadata: apievents.ServerMetadata{
ServerVersion: teleport.Version,
ServerID: s.uuid,
ServerNamespace: s.GetNamespace(),
},
Reason: events.SessionRejectedReasonMaxSessions,
Maximum: max,
}); err != nil {
s.logger.WarnContext(ctx, "Failed to emit session reject event", "error", err)
}
s.rejectChannel(ctx, nch, ssh.Prohibited, fmt.Sprintf("too many session channels for user %q (max=%d)", identityContext.TeleportUser, max))
return
}
decr = d
}
// SessionParams are not passed by old clients (<v19) or OpenSSH clients.
sessionParams, err := tracessh.ParseSessionParams(nch.ExtraData())
if err != nil {
s.logger.ErrorContext(ctx, "Failed to parse request data", "data", string(nch.ExtraData()), "error", err)
s.rejectChannel(ctx, nch, ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
return
}
ch, requests, err := nch.Accept()
if err != nil {
s.logger.WarnContext(ctx, "Unable to accept channel", "error", err)
s.rejectChannel(ctx, nch, ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
if decr != nil {
decr()
}
return
}
go func() {
s.handleSessionRequests(ctx, ccx, identityContext, sessionParams, ch, requests)
if decr != nil {
decr()
}
}()
// Channels of type "direct-tcpip" handles request for port forwarding.
case teleport.ChanDirectTCPIP:
// On regular server in "normal" mode "direct-tcpip" channels from
// SessionJoinPrincipal should be rejected, otherwise it's possible
// to use the "-teleport-internal-join" user to bypass RBAC.
if identityContext.Login == teleport.SSHSessionJoinPrincipal {
s.logger.ErrorContext(ctx, "Connection rejected, direct-tcpip with SessionJoinPrincipal in regular node must be blocked")
s.rejectChannel(
ctx,
nch, ssh.Prohibited,
fmt.Sprintf("attempted %v channel open in join-only mode", channelType))
return
}
req, err := sshutils.ParseDirectTCPIPReq(nch.ExtraData())
if err != nil {
s.logger.ErrorContext(ctx, "Failed to parse request data", "data", string(nch.ExtraData()), "error", err)
s.rejectChannel(ctx, nch, ssh.UnknownChannelType, "failed to parse direct-tcpip request")
return
}
ch, reqC, err := nch.Accept()
if err != nil {
s.logger.WarnContext(ctx, "Unable to accept channel", "error", err)
s.rejectChannel(ctx, nch, ssh.ConnectionFailed, fmt.Sprintf("unable to accept channel: %v", err))
return
}
go ssh.DiscardRequests(reqC)
go s.handleDirectTCPIPRequest(ctx, ccx, identityContext, ch, req)
default:
s.rejectChannel(ctx, nch, ssh.UnknownChannelType, fmt.Sprintf("unknown channel type: %v", channelType))
}
}
// canPortForward determines if port forwarding is allowed for the current
// user/role/node combo. Returns nil if port forwarding is allowed, non-nil
// if denied.
func (s *Server) canPortForward(scx *srv.ServerContext, mode decisionpb.SSHPortForwardMode) error {
// Is the node configured to allow port forwarding?
if !s.allowTCPForwarding {
return trace.AccessDenied("node does not allow port forwarding")
}
// Check if the role allows port forwarding for this user.
err := s.authHandlers.CheckPortForward(scx.DstAddr, scx, mode)
if err != nil {
return trace.Wrap(err)
}
return nil
}
// stderrWriter wraps an ssh.Channel in an implementation of io.StringWriter
// that sends anything written back the client over its stderr stream
type stderrWriter struct {
writer func(s string)
}
func (w *stderrWriter) WriteString(s string) (int, error) {
w.writer(s)
return len(s), nil
}
// handleDirectTCPIPRequest handles port forwarding requests.
func (s *Server) handleDirectTCPIPRequest(ctx context.Context, ccx *sshutils.ConnectionContext, identityContext srv.IdentityContext, channel ssh.Channel, req *sshutils.DirectTCPIPReq) {
// Create context for this channel. This context will be closed when
// forwarding is complete.
scx, err := srv.NewServerContext(ctx, ccx, s, identityContext, nil)
if err != nil {
s.logger.ErrorContext(ctx, "Unable to create connection context", "error", err)
s.writeStderr(ctx, channel, "Unable to create connection context.")
if err := channel.Close(); err != nil {
s.logger.WarnContext(ctx, "Failed to close channel", "error", err)
}
return
}
scx.IsTestStub = s.isTestStub
scx.TestLoginShell = s.testLoginShell
scx.AddCloser(channel)
scx.SessionRecordingConfig.SetMode(types.RecordOff)
scx.ExecType = teleport.ChanDirectTCPIP
scx.SrcAddr = sshutils.JoinHostPort(req.Orig, req.OrigPort)
scx.DstAddr = sshutils.JoinHostPort(req.Host, req.Port)
scx.SetAllowFileCopying(s.allowFileCopying)
defer scx.Close()
channel = scx.TrackActivity(channel)
// Bail out now if TCP port forwarding is not allowed for this node/user/role
// combo
if err = s.canPortForward(scx, decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_LOCAL); err != nil {
s.writeStderr(ctx, channel, err.Error())
return
}
scx.Logger.DebugContext(ctx, "Opening direct-tcpip channel", "source_addr", scx.SrcAddr, "dest_addr", scx.DstAddr)
defer scx.Logger.DebugContext(ctx, "Closing direct-tcpip channel", "source_addr", scx.SrcAddr, "dest_addr", scx.DstAddr)
conn, err := s.dialTCPIP(ctx, scx, scx.DstAddr)
if err != nil {
if errors.Is(err, trace.NotFound("%s", user.UnknownUserError(scx.Identity.Login))) || errors.Is(err, trace.BadParameter("unknown user")) {
// user does not exist for the provided login. Terminate the connection.
scx.Logger.WarnContext(ctx, "terminating direct-tcpip request because user does not exist", "user", scx.Identity.Login)
if err := ccx.ServerConn.Close(); err != nil {
scx.Logger.WarnContext(ctx, "Unable to terminate connection", "error", err)
}
return
}
scx.Logger.ErrorContext(ctx, "Forwarding data via direct-tcpip channel failed", "error", err)
s.writeStderr(ctx, channel, err.Error())
return
}
startEvent := scx.GetPortForwardEvent(events.PortForwardLocalEvent, events.PortForwardCode, scx.DstAddr)
s.emitAuditEventWithLog(ctx, &startEvent)
if err := utils.ProxyConn(ctx, conn, channel); err != nil && !errors.Is(err, io.EOF) && !errors.Is(err, os.ErrClosed) {
errEvent := scx.GetPortForwardEvent(events.PortForwardLocalEvent, events.PortForwardFailureCode, scx.DstAddr)
s.emitAuditEventWithLog(ctx, &errEvent)
scx.Logger.WarnContext(ctx, "Connection problem in direct-tcpip channel", "error", err)
}
stopEvent := scx.GetPortForwardEvent(events.PortForwardLocalEvent, events.PortForwardStopCode, scx.DstAddr)
s.emitAuditEventWithLog(ctx, &stopEvent)
}
// handleSessionRequests handles out of band session requests once the session
// channel has been created this function's loop handles all the "exec",
// "subsystem" and "shell" requests.
func (s *Server) handleSessionRequests(ctx context.Context, ccx *sshutils.ConnectionContext, identityContext srv.IdentityContext, sessionParams *tracessh.SessionParams, ch ssh.Channel, in <-chan *ssh.Request) {
netConfig, err := s.GetAccessPoint().GetClusterNetworkingConfig(ctx)
if err != nil {
s.logger.ErrorContext(ctx, "Unable to fetch cluster networking config", "error", err)
s.writeStderr(ctx, ch, "Unable to fetch cluster networking configuration.")
return
}
// Create context for this channel. This context will be closed when the
// session request is complete.
scx, err := srv.NewServerContext(ctx, ccx, s, identityContext, sessionParams, func(cfg *srv.MonitorConfig) {
cfg.IdleTimeoutMessage = netConfig.GetClientIdleTimeoutMessage()
cfg.MessageWriter = &stderrWriter{writer: func(msg string) { s.writeStderr(ctx, ch, msg) }}
})
if err != nil {
s.logger.ErrorContext(ctx, "Unable to create connection context", "error", err)
s.writeStderr(ctx, ch, "Unable to create connection context.")
if err := ch.Close(); err != nil {
s.logger.WarnContext(ctx, "Failed to close channel", "error", err)
}
return
}
scx.IsTestStub = s.isTestStub
scx.TestLoginShell = s.testLoginShell
scx.AddCloser(ch)
scx.ExecType = teleport.ChanSession
scx.SetAllowFileCopying(s.allowFileCopying)
defer scx.Close()
trackingChan := scx.TrackActivity(ch)
// If we are creating a new session (not joining a session), prepare a new session
// ID and inform the client.
//
// Note: If this is an old client (<v19), the join sid has not yet propagated
// from env vars. There is no harm in sending the ephemeral session ID anyways
// as clients should ignore the reported session ID when joining a session.
if scx.GetSessionParams().JoinSessionID == "" {
sid := session.NewID()
scx.SetNewSessionID(ctx, sid)
// inform the client of the session ID that is going to be used in a new
// goroutine to reduce latency.
go func() {
s.logger.DebugContext(ctx, "Sending current session ID", "sid", sid)
_, err := ch.SendRequest(teleport.CurrentSessionIDRequest, false, []byte(sid))
if err != nil {
s.logger.DebugContext(ctx, "Failed to send the current session ID", "error", err)
}
}()
}
// The keep-alive loop will keep pinging the remote server and after it has
// missed a certain number of keep-alive requests it will cancel the
// closeContext which signals the server to shutdown.
go srv.StartKeepAliveLoop(srv.KeepAliveParams{
Conns: []srv.RequestSender{
scx.ServerConn,
},
Interval: netConfig.GetKeepAliveInterval(),
MaxCount: netConfig.GetKeepAliveCountMax(),
CloseContext: ctx,
CloseCancel: scx.CancelFunc(),
})
for {
select {
case creq := <-scx.SubsystemResultCh:
// this means that subsystem has finished executing and
// want us to close session and the channel
scx.Logger.DebugContext(ctx, "Close session request", "error", creq.Err)
return
case req := <-in:
if req == nil {
// this will happen when the client closes/drops the connection
scx.Logger.DebugContext(ctx, "Client disconnected.", "client_addr", scx.ServerConn.RemoteAddr())
return
}
reqCtx := tracessh.ContextFromRequest(req)
ctx, span := s.tracerProvider.Tracer("ssh").Start(
oteltrace.ContextWithRemoteSpanContext(ctx, oteltrace.SpanContextFromContext(reqCtx)),
fmt.Sprintf("ssh.Regular.SessionRequest/%s", req.Type),
oteltrace.WithSpanKind(oteltrace.SpanKindServer),
oteltrace.WithAttributes(
semconv.RPCServiceKey.String("ssh.RegularServer"),
semconv.RPCMethodKey.String("SessionRequest"),
semconv.RPCSystemKey.String("ssh"),
),
)
// some functions called inside dispatch() may handle replies to SSH channel requests internally,
// rather than leaving the reply to be handled inside this loop. in that case, those functions must
// set req.WantReply to false so that two replies are not sent.
if err := s.dispatch(ctx, trackingChan, req, scx); err != nil {
s.replyError(ctx, trackingChan, req, err)
span.End()
return
}
if req.WantReply {
if err := req.Reply(true, nil); err != nil {
scx.Logger.WarnContext(ctx, "Failed to reply to request", "request_type", req.Type, "error", err)
}
}
span.End()
case result := <-scx.ExecResultCh:
scx.Logger.DebugContext(ctx, "Exec request complete", "command", result.Command, "code", result.Code)
// The exec process has finished and delivered the execution result, send
// the result back to the client, and close the session and channel.
_, err := trackingChan.SendRequest("exit-status", false, ssh.Marshal(struct{ C uint32 }{C: uint32(result.Code)}))
if err != nil {
scx.Logger.InfoContext(ctx, "Failed to send exit status", "command", result.Command, "error", err)
}
return
case <-ctx.Done():
scx.Logger.DebugContext(ctx, "Closing session due to cancellation")
return
}
}
}
// dispatch receives an SSH request for a subsystem and dispatches the request to the
// appropriate subsystem implementation
func (s *Server) dispatch(ctx context.Context, ch ssh.Channel, req *ssh.Request, serverContext *srv.ServerContext) error {
serverContext.Logger.DebugContext(ctx, "Handling session request", "request_type", req.Type, "want_reply", req.WantReply)
// If this SSH server is configured to only proxy, we do not support anything
// other than our own custom "subsystems" and environment manipulation.
if s.proxyMode {
switch req.Type {
case sshutils.SubsystemRequest:
return s.handleSubsystem(ctx, ch, req, serverContext)
case sshutils.EnvRequest:
return s.handleEnv(ctx, ch, req, serverContext)
case tracessh.EnvsRequest:
return s.handleEnvs(ctx, ch, req, serverContext)
case sshutils.AgentForwardRequest:
// process agent forwarding, but we will only forward agent to proxy in
// recording proxy mode.
err := s.handleAgentForwardProxy(ctx, serverContext)
if err != nil {
serverContext.Logger.WarnContext(ctx, "Failure forwarding agent", "error", err)
}
return nil
case sshutils.PuTTYSimpleRequest:
// PuTTY automatically requests a named 'simple@putty.projects.tartarus.org' channel any time it connects to a server
// as a proxy to indicate that it's in "simple" node and won't be requesting any other channels.
// As we don't support this request, we ignore it.
// https://the.earth.li/~sgtatham/putty/0.76/htmldoc/AppendixG.html#sshnames-channel
serverContext.Logger.DebugContext(ctx, "deliberately ignoring simple@putty.projects.tartarus.org request")
return nil
default:
s.logger.WarnContext(ctx, "server doesn't support request type", "request_type", req.Type)
if req.WantReply {
if err := req.Reply(false, nil); err != nil {
serverContext.Logger.ErrorContext(ctx, "error sending reply on SSH channel", "error", err)
}
}
return nil
}
}
// Certs with a join-only principal can only use a
// subset of all the possible request types.
if serverContext.JoinOnly {
switch req.Type {
case sshutils.PTYRequest:
return s.termHandlers.HandlePTYReq(ctx, ch, req, serverContext)
case sshutils.ShellRequest:
return s.termHandlers.HandleShell(ctx, ch, req, serverContext)
case sshutils.WindowChangeRequest:
return s.termHandlers.HandleWinChange(ctx, ch, req, serverContext)
case teleport.ForceTerminateRequest:
return s.termHandlers.HandleForceTerminate(ch, req, serverContext)
case sshutils.EnvRequest, tracessh.EnvsRequest:
// We ignore all SSH setenv requests for join-only principals.
// SSH will send them anyway but it seems fine to silently drop them.
case constants.InitiateFileTransfer:
if mode := serverContext.GetSessionParams().JoinMode; mode != types.SessionPeerMode {
return trace.AccessDenied("attempted file transfer in %s mode", mode)
}
return s.termHandlers.HandleFileTransferRequest(ctx, ch, req, serverContext)
case constants.FileTransferDecision:
return s.termHandlers.HandleFileTransferDecision(ctx, ch, req, serverContext)
case sshutils.SubsystemRequest:
return s.handleSubsystem(ctx, ch, req, serverContext)
case sshutils.AgentForwardRequest:
// This happens when SSH client has agent forwarding enabled, in this case
// client sends a special request, in return SSH server opens new channel
// that uses SSH protocol for agent drafted here:
// https://tools.ietf.org/html/draft-ietf-secsh-agent-02
// the open ssh proto spec that we implement is here:
// http://cvsweb.openbsd.org/cgi-bin/cvsweb/src/usr.bin/ssh/PROTOCOL.agent
// to maintain interoperability with OpenSSH, agent forwarding requests
// should never fail, all errors should be logged and we should continue
// processing requests.
err := s.handleAgentForwardNode(ctx, req, serverContext)
if err != nil {
serverContext.Logger.WarnContext(ctx, "failure forwarding agent", "error", err)
if trace.IsAccessDenied(err) {
s.writeStderr(ctx, ch, "Agent forwarding is not permitted for this user.\n")
} else {
s.writeStderr(ctx, ch, "Agent forwarding failed.\n")
}
}
return nil
case sshutils.PuTTYWinadjRequest:
return s.handlePuTTYWinadj(ctx, req)
default:
return trace.AccessDenied("attempted %v request in join-only mode", req.Type)
}
}
switch req.Type {
case sshutils.ExecRequest:
return s.termHandlers.HandleExec(ctx, ch, req, serverContext)
case sshutils.PTYRequest:
return s.termHandlers.HandlePTYReq(ctx, ch, req, serverContext)
case sshutils.ShellRequest:
return s.termHandlers.HandleShell(ctx, ch, req, serverContext)
case constants.InitiateFileTransfer:
return s.termHandlers.HandleFileTransferRequest(ctx, ch, req, serverContext)
case constants.FileTransferDecision:
return s.termHandlers.HandleFileTransferDecision(ctx, ch, req, serverContext)
case sshutils.WindowChangeRequest:
return s.termHandlers.HandleWinChange(ctx, ch, req, serverContext)
case teleport.ForceTerminateRequest:
return s.termHandlers.HandleForceTerminate(ch, req, serverContext)
case sshutils.EnvRequest:
return s.handleEnv(ctx, ch, req, serverContext)
case tracessh.EnvsRequest:
return s.handleEnvs(ctx, ch, req, serverContext)
case sshutils.SubsystemRequest:
// subsystems are SSH subsystems defined in http://tools.ietf.org/html/rfc4254 6.6
// they are in essence SSH session extensions, allowing to implement new SSH commands
return s.handleSubsystem(ctx, ch, req, serverContext)
case x11.ForwardRequest:
return s.handleX11Forward(ctx, ch, req, serverContext)
case sshutils.AgentForwardRequest:
// This happens when SSH client has agent forwarding enabled, in this case
// client sends a special request, in return SSH server opens new channel
// that uses SSH protocol for agent drafted here:
// https://tools.ietf.org/html/draft-ietf-secsh-agent-02
// the open ssh proto spec that we implement is here:
// http://cvsweb.openbsd.org/cgi-bin/cvsweb/src/usr.bin/ssh/PROTOCOL.agent
// to maintain interoperability with OpenSSH, agent forwarding requests
// should never fail, all errors should be logged and we should continue
// processing requests.
err := s.handleAgentForwardNode(ctx, req, serverContext)
if err != nil {
serverContext.Logger.WarnContext(ctx, "failure forwarding agent", "error", err)
if trace.IsAccessDenied(err) {
s.writeStderr(ctx, ch, "Agent forwarding is not permitted for this user.\n")
} else {
s.writeStderr(ctx, ch, "Agent forwarding failed.\n")
}
}
return nil
case sshutils.PuTTYWinadjRequest:
return s.handlePuTTYWinadj(ctx, req)
default:
serverContext.Logger.WarnContext(ctx, "server doesn't support request type", "request_type", req.Type)
if req.WantReply {
if err := req.Reply(false, nil); err != nil {
serverContext.Logger.ErrorContext(ctx, "error sending reply on SSH channel", "error", err)
}
}
return nil
}
}
// handleAgentForwardNode will create a unix socket and serve the agent running
// on the client on it.
func (s *Server) handleAgentForwardNode(ctx context.Context, _ *ssh.Request, scx *srv.ServerContext) (err error) {
event := scx.GetAgentForwardEvent()
defer func() {
if err != nil {
event.Metadata.Code = events.AgentForwardFailureCode
event.Status.Success = false
event.Status.Error = err.Error()
}
s.emitAuditEventWithLog(ctx, event)
}()
// check if the user's RBAC role allows agent forwarding
if err := s.authHandlers.CheckAgentForward(scx); err != nil {
return trace.Wrap(err)
}
// Enable agent forwarding for the broader connection-level
// context.
scx.Parent().SetForwardAgent(true)
if err := s.serveAgent(ctx, scx); err != nil {
return trace.Wrap(err)
}
return nil
}
// serveAgent will build the a sock path for this user and serve an SSH agent on unix socket.
func (s *Server) serveAgent(ctx context.Context, scx *srv.ServerContext) error {
proc, err := s.getNetworkingProcess(ctx, scx)
if err != nil {
return trace.Wrap(err)
}
listener, err := proc.ListenAgent(ctx)
if err != nil {
return trace.Wrap(err)
}
// start an agent server on a unix socket. each incoming connection
// will result in a separate agent request.
agentServer := sshagent.NewServer(func() (sshagent.Client, error) {
return scx.Parent().StartAgentChannel()
})
agentServer.SetListener(listener)
scx.Parent().AddCloser(agentServer)
scx.Parent().SetEnv(teleport.SSHAuthSock, listener.Addr().String())
scx.Parent().SetEnv(teleport.SSHAgentPID, fmt.Sprintf("%v", os.Getpid()))
scx.Logger.DebugContext(ctx, "Starting agent server for user", "teleport_user", scx.Identity.TeleportUser, "socket", agentServer.Path)
go func() {
if err := agentServer.Serve(); err != nil {
scx.Logger.ErrorContext(ctx, "agent server for user stopped", "teleport_user", scx.Identity.TeleportUser, "error", err)
}
}()
return nil
}
// handleAgentForwardProxy will forward the clients agent to the proxy (when
// the proxy is running in recording mode). When running in normal mode, this
// request will do nothing. To maintain interoperability, agent forwarding
// requests should never fail, all errors should be logged and we should
// continue processing requests.
func (s *Server) handleAgentForwardProxy(ctx context.Context, scx *srv.ServerContext) error {
// Forwarding an agent to the proxy is only supported when the proxy is in
// recording mode.
if !services.IsRecordAtProxy(scx.SessionRecordingConfig.GetMode()) {
return trace.BadParameter("agent forwarding to proxy only supported in recording mode")
}
if err := s.authHandlers.CheckAgentForward(scx); err != nil {
return trace.Wrap(err)
}
scx.Parent().SetForwardAgent(true)
return nil
}
// handleX11Forward handles an X11 forwarding request from the client.
func (s *Server) handleX11Forward(ctx context.Context, ch ssh.Channel, req *ssh.Request, scx *srv.ServerContext) (err error) {
event := &apievents.X11Forward{
Metadata: apievents.Metadata{
Type: events.X11ForwardEvent,
Code: events.X11ForwardCode,
},
UserMetadata: scx.Identity.GetUserMetadata(),
ConnectionMetadata: apievents.ConnectionMetadata{
LocalAddr: scx.ServerConn.LocalAddr().String(),
RemoteAddr: scx.ServerConn.RemoteAddr().String(),
},
ServerMetadata: scx.ServerMetadata(),
Status: apievents.Status{
Success: true,
},
}
defer func() {
if err != nil {
event.Metadata.Code = events.X11ForwardFailureCode
event.Status.Success = false
event.Status.Error = err.Error()
}
if trace.IsAccessDenied(err) {
// denied X11 requests are ok from a protocol perspective so we
// don't return them, just reply over ssh and emit the audit s.Logger.
s.replyError(ctx, ch, req, err)
err = nil
}
s.emitAuditEventWithLog(s.ctx, event)
}()
// check if X11 forwarding is disabled, or if xauth can't be handled.
if !s.x11.Enabled || x11.CheckXAuthPath() != nil {
return trace.AccessDenied("X11 forwarding is not enabled")
}
// Check if the user's RBAC role allows X11 forwarding.
if err := s.authHandlers.CheckX11Forward(scx); err != nil {
return trace.Wrap(err)
}
var x11Req x11.ForwardRequestPayload
if err := ssh.Unmarshal(req.Payload, &x11Req); err != nil {
return trace.Wrap(err)
}
proc, err := s.getNetworkingProcess(ctx, scx)
if err != nil {
return trace.Wrap(err)
}
listener, err := proc.ListenX11(ctx, networking.X11Request{
ForwardRequestPayload: x11Req,
DisplayOffset: s.x11.DisplayOffset,
MaxDisplay: s.x11.MaxDisplay,
})
if err != nil {
return trace.Wrap(err)
}
scx.Parent().AddCloser(listener)
if err := scx.HandleX11Listener(ctx, listener, x11Req.SingleConnection); err != nil {
if trace.IsLimitExceeded(err) {
return trace.AccessDenied("The server cannot support any more X11 forwarding sessions at this time")
}
return trace.Wrap(err)
}
return nil
}
func (s *Server) handleSubsystem(ctx context.Context, ch ssh.Channel, req *ssh.Request, serverContext *srv.ServerContext) error {
sb, err := s.parseSubsystemRequest(ctx, req, serverContext)
if err != nil {
serverContext.Logger.WarnContext(ctx, "Failed to parse subsystem request", "request_type", req.Type, "error", err)
return trace.Wrap(err)
}
serverContext.Logger.DebugContext(ctx, "Starting subsystem")
// starting subsystem is blocking to the client,
// while collecting its result and waiting is not blocking
if err := sb.Start(ctx, serverContext.ServerConn, ch, req, serverContext); err != nil {
serverContext.Logger.WarnContext(ctx, "Subsystem request failed", "error", err)
serverContext.SendSubsystemResult(ctx, srv.SubsystemResult{Err: trace.Wrap(err)})
return trace.Wrap(err)
}
go func() {
err := sb.Wait()
serverContext.Logger.DebugContext(ctx, "Subsystem finished", "subsystem", sb, "error", err)
serverContext.SendSubsystemResult(ctx, srv.SubsystemResult{Err: trace.Wrap(err)})
}()
return nil
}
// handleEnv accepts an environment variable sent by the client and stores it
// in connection context
func (s *Server) handleEnv(ctx context.Context, ch ssh.Channel, req *ssh.Request, scx *srv.ServerContext) error {
var e sshutils.EnvReqParams
if err := ssh.Unmarshal(req.Payload, &e); err != nil {
scx.Logger.ErrorContext(ctx, "failed to parse env request", "error", err)
return trace.Wrap(err, "failed to parse env request")
}
scx.SetEnv(e.Name, e.Value)
return nil
}
// handleEnvs accepts environment variables sent by the client and stores them
// in connection context
func (s *Server) handleEnvs(ctx context.Context, ch ssh.Channel, req *ssh.Request, scx *srv.ServerContext) error {
var raw tracessh.EnvsReq
if err := ssh.Unmarshal(req.Payload, &raw); err != nil {
scx.Logger.ErrorContext(ctx, "failed to parse envs request", "error", err)
return trace.Wrap(err, "failed to parse envs request")
}
var envs map[string]string
if err := json.Unmarshal(raw.EnvsJSON, &envs); err != nil {
return trace.Wrap(err, "failed to unmarshal envs")
}
for k, v := range envs {
scx.SetEnv(k, v)
}
return nil
}
// handleKeepAlive accepts and replies to keepalive@openssh.com requests.
func (s *Server) handleKeepAlive(req *ssh.Request) {
// only reply if the sender actually wants a response
if !req.WantReply {
return
}
if err := req.Reply(true, nil); err != nil {
s.logger.WarnContext(s.ctx, "Unable to reply to request", "request_type", req.Type, "error", err)
return
}
s.logger.DebugContext(s.ctx, "successfully replied to request", "request_type", req.Type)
}
// handleClusterDetails responds to global out-of-band with details about the cluster.
func (s *Server) handleClusterDetails(ctx context.Context, req *ssh.Request) {
s.logger.DebugContext(ctx, "cluster details request received")
if !req.WantReply {
return
}
// get the cluster config, if we can't get it, reply false
recConfig, err := s.authService.GetSessionRecordingConfig(ctx)
if err != nil {
if err := req.Reply(false, nil); err != nil {
s.logger.WarnContext(ctx, "Unable to respond to cluster details request", "error", err)
}
return
}
details := sshutils.ClusterDetails{
RecordingProxy: services.IsRecordAtProxy(recConfig.GetMode()),
FIPSEnabled: s.fips,
}
if err = req.Reply(true, ssh.Marshal(details)); err != nil {
s.logger.WarnContext(ctx, "Unable to respond to cluster details request", "error", err)
return
}
s.logger.DebugContext(ctx, "Replied to cluster details request")
}
// handleVersionRequest replies with the Teleport version of the server.
func (s *Server) handleVersionRequest(ctx context.Context, req *ssh.Request) {
err := req.Reply(true, []byte(teleport.Version))
if err != nil {
s.logger.DebugContext(ctx, "Failed to reply to version request", "error", err)
}
}
// handleProxyJump handles ProxyJump request that is executed via direct tcp-ip dial on the proxy
func (s *Server) handleProxyJump(ctx context.Context, ccx *sshutils.ConnectionContext, identityContext srv.IdentityContext, ch ssh.Channel, req sshutils.DirectTCPIPReq) {
// Create context for this channel. This context will be closed when the
// session request is complete.
scx, err := srv.NewServerContext(ctx, ccx, s, identityContext, nil)
if err != nil {
s.logger.ErrorContext(ctx, "Unable to create connection context", "error", err)
s.writeStderr(ctx, ch, "Unable to create connection context.")
if err := ch.Close(); err != nil {
s.logger.WarnContext(ctx, "Failed to close channel", "error", err)
}
return
}
scx.IsTestStub = s.isTestStub
scx.TestLoginShell = s.testLoginShell
scx.AddCloser(ch)
scx.SetAllowFileCopying(s.allowFileCopying)
defer scx.Close()
ch = scx.TrackActivity(ch)
recConfig, err := s.GetAccessPoint().GetSessionRecordingConfig(ctx)
if err != nil {
s.logger.ErrorContext(ctx, "Unable to fetch session recording config", "error", err)
s.writeStderr(ctx, ch, "Unable to fetch session recording configuration.")
return
}
// force agent forward, because in recording mode proxy needs
// client's agent to authenticate to the target server
//
// When proxy is in "Recording mode" the following will happen with SSH:
//
// $ ssh -J user@teleport.proxy:3023 -p 3022 user@target -F ./forward.config
//
// Where forward.config enables agent forwarding:
//
// Host teleport.proxy
// ForwardAgent yes
//
// This will translate to ProxyCommand:
//
// exec ssh -l user -p 3023 -F ./forward.config -vvv -W 'target:3022' teleport.proxy
//
// -W means establish direct tcp-ip, and in SSH 2.0 session implementation,
// this gets called before agent forwarding is requested:
//
// https://github.com/openssh/openssh-portable/blob/master/ssh.c#L1884
//
// so in recording mode, proxy is forced to request agent forwarding
// "out of band", before SSH client actually asks for it
// which is a hack, but the only way we can think of making it work,
// ideas are appreciated.
if services.IsRecordAtProxy(recConfig.GetMode()) {
err = s.handleAgentForwardProxy(ctx, scx)
if err != nil {
s.logger.WarnContext(ctx, "Failed to request agent in recording mode", "error", err)
s.writeStderr(ctx, ch, "Failed to request agent")
return
}
}
netConfig, err := s.GetAccessPoint().GetClusterNetworkingConfig(ctx)
if err != nil {
s.logger.ErrorContext(ctx, "Unable to fetch cluster networking config", "error", err)
s.writeStderr(ctx, ch, "Unable to fetch cluster networking configuration.")
return
}
// The keep-alive loop will keep pinging the remote server and after it has
// missed a certain number of keep-alive requests it will cancel the
// closeContext which signals the server to shutdown.
go srv.StartKeepAliveLoop(srv.KeepAliveParams{
Conns: []srv.RequestSender{
scx.ServerConn,
},
Interval: netConfig.GetKeepAliveInterval(),
MaxCount: netConfig.GetKeepAliveCountMax(),
CloseContext: ctx,
CloseCancel: scx.CancelFunc(),
})
subsys, err := newProxySubsys(ctx, scx, s, proxySubsysRequest{
host: req.Host,
port: fmt.Sprintf("%v", req.Port),
})
if err != nil {
s.logger.ErrorContext(ctx, "Unable instantiate proxy subsystem", "error", err)
s.writeStderr(ctx, ch, "Unable to instantiate proxy subsystem.")
return
}
if err := subsys.Start(ctx, scx.ServerConn, ch, &ssh.Request{}, scx); err != nil {
s.logger.ErrorContext(ctx, "Unable to start proxy subsystem", "error", err)
s.writeStderr(ctx, ch, "Unable to start proxy subsystem.")
return
}
wch := make(chan struct{})
go func() {
defer close(wch)
if err := subsys.Wait(); err != nil {
s.logger.ErrorContext(ctx, "Proxy subsystem failed", "error", err)
s.writeStderr(ctx, ch, "Proxy subsystem failed.")
}
}()
select {
case <-wch:
case <-ctx.Done():
}
}
// createForwardingContext creates a server context for a user intending to
// port forward. It returns an error if the user is not allowed to port forward.
func (s *Server) createForwardingContext(ctx context.Context, ccx *sshutils.ConnectionContext, r *ssh.Request) (context.Context, *srv.ServerContext, error) {
req, err := sshutils.ParseTCPIPForwardReq(r.Payload)
if err != nil {
return nil, nil, trace.Wrap(err)
}
identityContext, err := s.authHandlers.CreateIdentityContext(ccx.ServerConn)
if err != nil {
return nil, nil, trace.Wrap(err)
}
// On regular server in "normal" mode "tcpip-forward" requests from
// SessionJoinPrincipal should be rejected, otherwise it's possible to use
// the "-teleport-internal-join" user to bypass RBAC.
if identityContext.Login == teleport.SSHSessionJoinPrincipal {
s.logger.ErrorContext(ctx, "Request with SessionJoinPrincipal rejected", "request_type", r.Type)
err := trace.AccessDenied("attempted %q request in join-only mode", r.Type)
if replyErr := r.Reply(false, []byte(utils.FormatErrorWithNewline(err))); replyErr != nil {
s.logger.WarnContext(ctx, "Failed to reply to request", "request_type", r.Type, "error", err)
}
// Disable default reply by caller, we already handled it.
r.WantReply = false
return nil, nil, err
}
// Create context for this request.
scx, err := srv.NewServerContext(ctx, ccx, s, identityContext, nil)
if err != nil {
return nil, nil, trace.Wrap(err)
}
listenAddr := sshutils.JoinHostPort(req.Addr, req.Port)
scx.IsTestStub = s.isTestStub
scx.TestLoginShell = s.testLoginShell
scx.ExecType = teleport.TCPIPForwardRequest
scx.SrcAddr = listenAddr
scx.DstAddr = ccx.NetConn.RemoteAddr().String()
scx.SessionRecordingConfig.SetMode(types.RecordOff)
scx.SetAllowFileCopying(s.allowFileCopying)
if err := s.canPortForward(scx, decisionpb.SSHPortForwardMode_SSH_PORT_FORWARD_MODE_REMOTE); err != nil {
scx.Close()
return nil, nil, trace.Wrap(err)
}
return ctx, scx, nil
}
// handleTCPIPForwardRequest handles remote port forwarding requests.
func (s *Server) handleTCPIPForwardRequest(ctx context.Context, ccx *sshutils.ConnectionContext, r *ssh.Request) error {
ctx, scx, err := s.createForwardingContext(ctx, ccx, r)
if err != nil {
return trace.Wrap(err)
}
listener, err := s.listenTCPIP(ctx, scx, scx.SrcAddr)
if err != nil {
if serr := scx.Close(); serr != nil {
s.logger.DebugContext(ctx, "Failed while cleaning up request",
"request_type", teleport.TCPIPForwardRequest,
"server_context_close_error", serr,
"error", err)
}
return trace.Wrap(err)
}
// If the client didn't request a specific port, the chosen port needs to
// be reported back.
srcHost, srcPort, err := sshutils.SplitHostPort(scx.SrcAddr)
if err != nil {
if lerr := listener.Close(); lerr != nil {
s.logger.DebugContext(ctx, "Failed while cleaning up request",
"request_type", teleport.TCPIPForwardRequest,
"listener_close_error", lerr,
"error", err)
}
if serr := scx.Close(); serr != nil {
s.logger.DebugContext(ctx, "Failed while cleaning up request",
"request_type", teleport.TCPIPForwardRequest,
"server_context_close_error", serr,
"error", err)
}
return trace.Wrap(err)
}
_, listenPort, err := sshutils.SplitHostPort(listener.Addr().String())
if err != nil {
if lerr := listener.Close(); lerr != nil {
s.logger.DebugContext(ctx, "Failed while cleaning up request",
"request_type", teleport.TCPIPForwardRequest,
"listener_close_error", lerr,
"error", err)
}
if serr := scx.Close(); serr != nil {
s.logger.DebugContext(ctx, "Failed while cleaning up request",
"request_type", teleport.TCPIPForwardRequest,
"server_context_close_error", serr,
"error", err)
}
return trace.Wrap(err)
}
scx.SrcAddr = sshutils.JoinHostPort(srcHost, listenPort)
event := scx.GetPortForwardEvent(events.PortForwardRemoteEvent, events.PortForwardCode, scx.SrcAddr)
s.emitAuditEventWithLog(ctx, &event)
// spawn remote forwarding handler to multiplex connections to the forwarded port
go func() {
defer scx.Close()
stopEvent := scx.GetPortForwardEvent(events.PortForwardRemoteEvent, events.PortForwardStopCode, scx.SrcAddr)
defer s.emitAuditEventWithLog(ctx, &stopEvent)
for {
conn, err := listener.Accept()
if err != nil {
if !utils.IsOKNetworkError(err) {
slog.WarnContext(ctx, "failed to accept connection", "error", err)
}
return
}
logger := slog.With(
"src_addr", scx.SrcAddr,
"remote_addr", conn.RemoteAddr().String(),
)
dstHost, dstPort, err := sshutils.SplitHostPort(conn.RemoteAddr().String())
if err != nil {
conn.Close()
logger.WarnContext(ctx, "failed to parse addr", "error", err)
return
}
req := sshutils.ForwardedTCPIPRequest{
Addr: srcHost,
Port: listenPort,
OrigAddr: dstHost,
OrigPort: dstPort,
}
if err := req.CheckAndSetDefaults(); err != nil {
conn.Close()
logger.WarnContext(ctx, "failed to create forwarded tcpip request", "error", err)
return
}
reqBytes := ssh.Marshal(req)
ch, rch, err := scx.ConnectionContext.ServerConn.OpenChannel(teleport.ChanForwardedTCPIP, reqBytes)
if err != nil {
conn.Close()
logger.WarnContext(ctx, "failed to open channel", "error", err)
continue
}
ch = scx.TrackActivity(ch)
go ssh.DiscardRequests(rch)
go io.Copy(io.Discard, ch.Stderr())
go func() {
startEvent := scx.GetPortForwardEvent(events.PortForwardRemoteConnEvent, events.PortForwardCode, scx.SrcAddr)
startEvent.RemoteAddr = conn.RemoteAddr().String()
s.emitAuditEventWithLog(ctx, &startEvent)
if err := utils.ProxyConn(ctx, conn, ch); err != nil {
errEvent := scx.GetPortForwardEvent(events.PortForwardRemoteConnEvent, events.PortForwardFailureCode, scx.SrcAddr)
errEvent.RemoteAddr = conn.RemoteAddr().String()
s.emitAuditEventWithLog(ctx, &errEvent)
}
stopEvent := scx.GetPortForwardEvent(events.PortForwardRemoteConnEvent, events.PortForwardStopCode, scx.SrcAddr)
stopEvent.RemoteAddr = conn.RemoteAddr().String()
s.emitAuditEventWithLog(ctx, &stopEvent)
}()
}
}()
// Report addr back to the client.
if r.WantReply {
var payload []byte
if srcPort == 0 {
payload = ssh.Marshal(struct {
Port uint32
}{Port: uint32(listener.Addr().(*net.TCPAddr).Port)})
}
if err := r.Reply(true, payload); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to request", "request_type", r.Type, "error", err)
}
}
key := getRemoteForwardingMapKey(scx)
s.remoteForwardingMap.Store(key, listener)
// Close the listener once the connection is closed, if it hasn't
// been closed already via a cancel-tcpip-forward request.
ccx.AddCloser(utils.CloseFunc(func() error {
listener, ok := s.remoteForwardingMap.LoadAndDelete(key)
if ok {
return trace.Wrap(listener.Close())
}
return nil
}))
return nil
}
// handleCancelTCPIPForwardRequest handles canceling a previously requested
// remote forwarded port.
func (s *Server) handleCancelTCPIPForwardRequest(ctx context.Context, ccx *sshutils.ConnectionContext, r *ssh.Request) error {
_, scx, err := s.createForwardingContext(ctx, ccx, r)
if err != nil {
return trace.Wrap(err)
}
defer scx.Close()
listener, ok := s.remoteForwardingMap.LoadAndDelete(getRemoteForwardingMapKey(scx))
if !ok {
return trace.NotFound("no remote forwarding listener at %v", scx.SrcAddr)
}
if err := r.Reply(true, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to request", "request_type", r.Type, "error", err)
}
return trace.Wrap(listener.Close())
}
func (s *Server) replyError(ctx context.Context, ch ssh.Channel, req *ssh.Request, err error) {
s.logger.ErrorContext(ctx, "failure handling SSH request", "request_type", req.Type, "error", err)
// Terminate the error with a newline when writing to remote channel's
// stderr so the output does not mix with the rest of the output if the remote
// side is not doing additional formatting for extended data.
// See github.com/gravitational/teleport/issues/4542
message := utils.FormatErrorWithNewline(err)
s.writeStderr(ctx, ch, message)
if req.WantReply {
if err := req.Reply(false, []byte(message)); err != nil {
s.logger.WarnContext(ctx, "Failed to reply with error to request", "request_type", req.Type, "error", err)
}
}
}
func (s *Server) parseSubsystemRequest(ctx context.Context, req *ssh.Request, serverContext *srv.ServerContext) (srv.Subsystem, error) {
var r sshutils.SubsystemReq
if err := ssh.Unmarshal(req.Payload, &r); err != nil {
return nil, trace.BadParameter("failed to parse subsystem request: %v", err)
}
if s.proxyMode {
switch {
case strings.HasPrefix(r.Name, "proxy:"):
return s.parseProxySubsys(ctx, r.Name, serverContext)
case strings.HasPrefix(r.Name, "proxysites"):
return parseProxySitesSubsys(r.Name, s)
default:
return nil, trace.BadParameter("unrecognized subsystem: %v", r.Name)
}
}
switch r.Name {
case teleport.SFTPSubsystem:
err := serverContext.CheckSFTPAllowed(s.reg)
if err != nil {
s.emitAuditEventWithLog(context.Background(), &apievents.SFTP{
Metadata: apievents.Metadata{
Code: events.SFTPDisallowedCode,
Type: events.SFTPEvent,
Time: time.Now(),
},
UserMetadata: serverContext.Identity.GetUserMetadata(),
ServerMetadata: serverContext.GetServer().EventMetadata(),
Error: err.Error(),
})
return nil, trace.Wrap(err)
}
return newSFTPSubsys(serverContext.ConsumeApprovedFileTransferRequest())
default:
return nil, trace.BadParameter("unrecognized subsystem: %v", r.Name)
}
}
func (s *Server) writeStderr(ctx context.Context, ch ssh.Channel, msg string) {
if _, err := io.WriteString(ch.Stderr(), msg); err != nil {
s.logger.WarnContext(ctx, "Failed writing to stderr of SSH channel", "error", err)
}
}
func (s *Server) rejectChannel(ctx context.Context, ch ssh.NewChannel, reason ssh.RejectionReason, msg string) {
if err := ch.Reject(reason, msg); err != nil {
s.logger.WarnContext(ctx, "Failed to reject new SSH channel", "error", err)
}
}
// handlePuTTYWinadj replies with failure to a PuTTY winadj request as required.
// it returns an error if the reply fails. context from the PuTTY documentation:
// PuTTY sends this request along with some SSH_MSG_CHANNEL_WINDOW_ADJUST messages as part of its window-size
// tuning. It can be sent on any type of channel. There is no message-specific data. Servers MUST treat it
// as an unrecognized request and respond with SSH_MSG_CHANNEL_FAILURE.
// https://the.earth.li/~sgtatham/putty/0.76/htmldoc/AppendixG.html#sshnames-channel
func (s *Server) handlePuTTYWinadj(ctx context.Context, req *ssh.Request) error {
if err := req.Reply(false, nil); err != nil {
s.logger.WarnContext(ctx, "Failed to reply to PuTTY winadj request", "error", err)
return err
}
// the reply has been handled inside this function (rather than relying on the standard behavior
// of leaving handleSessionRequests to do it) so set the WantReply flag to false here.
req.WantReply = false
return nil
}
func (s *Server) emitAuditEventWithLog(ctx context.Context, event apievents.AuditEvent) {
if err := s.EmitAuditEvent(ctx, event); err != nil {
s.logger.WarnContext(ctx, "Failed to emit event", "type", event.GetType(), "code", event.GetCode())
}
}
/*
* 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 utils
import (
"context"
"fmt"
"log/slog"
"net"
"net/url"
"strconv"
"strings"
"unicode/utf8"
"github.com/gravitational/trace"
apiutils "github.com/gravitational/teleport/api/utils"
)
// NetAddr is network address that includes network, optional path and
// host port
type NetAddr struct {
// Addr is the host:port address, like "localhost:22"
Addr string `json:"addr"`
// AddrNetwork is the type of a network socket, like "tcp" or "unix"
AddrNetwork string `json:"network,omitempty"`
// Path is a socket file path, like '/var/path/to/socket' in "unix:///var/path/to/socket"
Path string `json:"path,omitempty"`
}
// Host returns host part of address without port
func (a *NetAddr) Host() string {
host, _, err := net.SplitHostPort(a.Addr)
if err == nil {
return host
}
// this is done to remove optional square brackets
if ip := net.ParseIP(strings.Trim(a.Addr, "[]")); len(ip) != 0 {
return ip.String()
}
return a.Addr
}
// Port returns defaultPort if no port is set or is invalid,
// the real port otherwise
func (a *NetAddr) Port(defaultPort int) int {
_, port, err := net.SplitHostPort(a.Addr)
if err != nil {
return defaultPort
}
porti, err := strconv.Atoi(port)
if err != nil {
return defaultPort
}
return porti
}
// IsLocal returns true if this is a local address
func (a *NetAddr) IsLocal() bool {
host, _, err := net.SplitHostPort(a.Addr)
if err != nil {
return false
}
return IsLocalhost(host)
}
// IsLoopback returns true if this is a loopback address
func (a *NetAddr) IsLoopback() bool {
return apiutils.IsLoopback(a.Addr)
}
// IsHostUnspecified returns true if this address' host is unspecified.
func (a *NetAddr) IsHostUnspecified() bool {
return a.Host() == "" || net.ParseIP(a.Host()).IsUnspecified()
}
// IsEmpty returns true if address is empty
func (a *NetAddr) IsEmpty() bool {
return a == nil || (a.Addr == "" && a.AddrNetwork == "" && a.Path == "")
}
// FullAddress returns full address including network and address (tcp://0.0.0.0:1243)
func (a *NetAddr) FullAddress() string {
return fmt.Sprintf("%v://%v", a.AddrNetwork, a.Addr)
}
// String returns address without network (0.0.0.0:1234)
func (a *NetAddr) String() string {
return a.Addr
}
// Network returns the scheme for this network address (tcp or unix)
func (a *NetAddr) Network() string {
return a.AddrNetwork
}
// MarshalYAML defines how a network address should be marshaled to a string
func (a *NetAddr) MarshalYAML() (any, error) {
url := url.URL{Scheme: a.AddrNetwork, Host: a.Addr, Path: a.Path}
return strings.TrimLeft(url.String(), "/"), nil
}
// UnmarshalYAML defines how a string can be unmarshalled into a network address
func (a *NetAddr) UnmarshalYAML(unmarshal func(any) error) error {
var addr string
err := unmarshal(&addr)
if err != nil {
return err
}
parsedAddr, err := ParseAddr(addr)
if err != nil {
return err
}
*a = *parsedAddr
return nil
}
func (a *NetAddr) Set(s string) error {
v, err := ParseAddr(s)
if err != nil {
return trace.Wrap(err)
}
a.Addr = v.Addr
a.AddrNetwork = v.AddrNetwork
return nil
}
// NetAddrsToStrings takes a list of netAddrs and returns a list of address strings.
func NetAddrsToStrings(netAddrs []NetAddr) []string {
addrs := make([]string, len(netAddrs))
for i, addr := range netAddrs {
addrs[i] = addr.String()
}
return addrs
}
// ParseAddrs parses the provided slice of strings as a slice of NetAddr's.
func ParseAddrs(addrs []string) (result []NetAddr, err error) {
for _, addr := range addrs {
parsed, err := ParseAddr(addr)
if err != nil {
return nil, trace.Wrap(err)
}
result = append(result, *parsed)
}
return result, nil
}
// ParseAddr takes strings like "tcp://host:port/path" and returns
// *NetAddr or an error
func ParseAddr(a string) (*NetAddr, error) {
if a == "" {
return nil, trace.BadParameter("missing parameter address")
}
if !strings.Contains(a, "://") {
a = "tcp://" + a
}
u, err := url.Parse(a)
if err != nil {
return nil, trace.BadParameter("failed to parse %q: %v", a, err)
}
switch u.Scheme {
case "tcp":
return &NetAddr{Addr: u.Host, AddrNetwork: u.Scheme, Path: u.Path}, nil
case "unix":
return &NetAddr{Addr: u.Path, AddrNetwork: u.Scheme}, nil
case "http", "https":
return &NetAddr{Addr: u.Host, AddrNetwork: u.Scheme, Path: u.Path}, nil
default:
return nil, trace.BadParameter("%q: unsupported scheme: %q", a, u.Scheme)
}
}
// MustParseAddr parses the provided string into NetAddr or panics on an error
func MustParseAddr(a string) *NetAddr {
addr, err := ParseAddr(a)
if err != nil {
panic(fmt.Sprintf("failed to parse %v: %v", a, err))
}
return addr
}
// MustParseAddrList parses the provided list of strings into a NetAddr list or panics on error
func MustParseAddrList(aList ...string) []NetAddr {
addrList := make([]NetAddr, len(aList))
for i, a := range aList {
addrList[i] = *MustParseAddr(a)
}
return addrList
}
// FromAddr returns NetAddr from golang standard net.Addr
func FromAddr(a net.Addr) NetAddr {
return NetAddr{AddrNetwork: a.Network(), Addr: a.String()}
}
// JoinAddrSlices joins two addr slices and returns a resulting slice
func JoinAddrSlices(a []NetAddr, b []NetAddr) []NetAddr {
if len(a)+len(b) == 0 {
return nil
}
out := make([]NetAddr, 0, len(a)+len(b))
out = append(out, a...)
out = append(out, b...)
return out
}
// ParseHostPortAddr takes strings like "host:port" and returns
// *NetAddr or an error
//
// If defaultPort == -1 it expects 'hostport' string to have it
func ParseHostPortAddr(hostport string, defaultPort int) (*NetAddr, error) {
addr, err := ParseAddr(hostport)
if err != nil {
return nil, trace.Wrap(err)
}
// port is required but not set
if defaultPort == -1 && addr.Addr == addr.Host() {
return nil, trace.BadParameter("missing port in address %q", hostport)
}
addr.Addr = net.JoinHostPort(addr.Host(), fmt.Sprintf("%v", addr.Port(defaultPort)))
return addr, nil
}
// DialAddrFromListenAddr returns dial address from listen address
func DialAddrFromListenAddr(listenAddr NetAddr) NetAddr {
if listenAddr.IsEmpty() {
return listenAddr
}
return NetAddr{Addr: ReplaceLocalhost(listenAddr.Addr, "127.0.0.1")}
}
// ReplaceLocalhost checks if a given address is link-local (like 0.0.0.0 or 127.0.0.1)
// and replaces it with the IP taken from replaceWith, preserving the original port
//
// Both addresses are in "host:port" format
// The function returns the original value if it encounters any problems with parsing
func ReplaceLocalhost(addr, replaceWith string) string {
host, port, err := net.SplitHostPort(addr)
if err != nil {
return addr
}
if IsLocalhost(host) {
host, _, err = net.SplitHostPort(replaceWith)
if err != nil {
return addr
}
addr = net.JoinHostPort(host, port)
}
return addr
}
// IsLocalhost returns true if this is a local hostname or ip
func IsLocalhost(host string) bool {
if host == "localhost" {
return true
}
ip := net.ParseIP(host)
return ip.IsLoopback() || ip.IsUnspecified()
}
// GuessIP tries to guess an IP address this machine is reachable at on the
// internal network, always picking IPv4 from the internal address space
//
// If no internal IPs are found, it returns 127.0.0.1 but it never returns
// an address from the public IP space
func GuessHostIP() (ip net.IP, err error) {
ifaces, err := net.Interfaces()
if err != nil {
return nil, trace.Wrap(err)
}
adrs := make([]net.Addr, 0)
for _, iface := range ifaces {
ifadrs, err := iface.Addrs()
if err != nil {
slog.WarnContext(context.Background(), "Unable to get addresses for interface", "interface", iface.Name, "error", err)
} else {
adrs = append(adrs, ifadrs...)
}
}
return guessHostIP(adrs), nil
}
func guessHostIP(addrs []net.Addr) (ip net.IP) {
// collect the list of all IPv4s
var ips []net.IP
for _, addr := range addrs {
var ipAddr net.IP
a, ok := addr.(*net.IPAddr)
if ok {
ipAddr = a.IP
} else {
in, ok := addr.(*net.IPNet)
if ok {
ipAddr = in.IP
} else {
continue
}
}
if ipAddr.To4() == nil || ipAddr.IsLoopback() || ipAddr.IsMulticast() {
continue
}
ips = append(ips, ipAddr)
}
for i := range ips {
first := &net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}
second := &net.IPNet{IP: net.IPv4(192, 168, 0, 0), Mask: net.CIDRMask(16, 32)}
third := &net.IPNet{IP: net.IPv4(172, 16, 0, 0), Mask: net.CIDRMask(12, 32)}
// our first pick would be "10.0.0.0/8"
if first.Contains(ips[i]) {
ip = ips[i]
break
// our 2nd pick would be "192.168.0.0/16"
} else if second.Contains(ips[i]) {
ip = ips[i]
// our 3rd pick would be "172.16.0.0/12"
} else if third.Contains(ips[i]) && !second.Contains(ip) {
ip = ips[i]
}
}
if ip == nil {
if len(ips) > 0 {
return ips[0]
}
// fallback to loopback
ip = net.IPv4(127, 0, 0, 1)
}
return ip
}
// ReplaceUnspecifiedHost replaces unspecified "0.0.0.0" with localhost since "0.0.0.0" is never a valid
// principal (auth server explicitly removes it when issuing host certs) and when a reverse tunnel client used
// establishes SSH reverse tunnel connection the host is validated against
// the valid principal list.
func ReplaceUnspecifiedHost(addr *NetAddr, defaultPort int) string {
if !addr.IsHostUnspecified() {
return addr.String()
}
port := addr.Port(defaultPort)
return net.JoinHostPort("localhost", strconv.Itoa(port))
}
// ToLowerCaseASCII returns a lower-case version of in. See RFC 6125 6.4.1. We use
// an explicitly ASCII function to avoid any sharp corners resulting from
// performing Unicode operations on DNS labels.
//
// NOTE: copied verbatim from crypto/x509 source, including the above comments. Teleport
// uses this function to approximate a form of opt-in case-insensitivity for ssh hostnames
func ToLowerCaseASCII(in string) string {
// If the string is already lower-case then there's nothing to do.
isAlreadyLowerCase := true
for _, c := range in {
if c == utf8.RuneError {
// If we get a UTF-8 error then there might be
// upper-case ASCII bytes in the invalid sequence.
isAlreadyLowerCase = false
break
}
if 'A' <= c && c <= 'Z' {
isAlreadyLowerCase = false
break
}
}
if isAlreadyLowerCase {
return in
}
out := []byte(in)
for i, c := range out {
if 'A' <= c && c <= 'Z' {
out[i] += 'a' - 'A'
}
}
return string(out)
}
/*
* Teleport
* Copyright (C) 2024 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 utils
// Combinations yields all unique sub-slices of the input slice.
func Combinations(verbs []string) [][]string {
var result [][]string
length := len(verbs)
for i := range 1 << length {
subslice := make([]string, 0)
for j := range length {
if i&(1<<j) != 0 {
subslice = append(subslice, verbs[j])
}
}
result = append(result, subslice)
}
return result
}
/*
* 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 utils
import (
"crypto/hmac"
"crypto/sha256"
"encoding/base64"
"strings"
"github.com/gravitational/trace"
)
// Anonymizer defines an interface for anonymizing data
type Anonymizer interface {
// Anonymize returns anonymized string from the provided data
Anonymize(data []byte) string
// AnonymizeString anonymizes the given string data using HMAC
AnonymizeString(s string) string
// AnonymizeNonEmpty anonymizes the given string into bytes if the string is
// nonempty, otherwise returns an empty slice.
AnonymizeNonEmpty(s string) []byte
}
var _ AnonymizationKeyProvider = (AnonymizationKeyString)("")
// AnonymizationKeyString is a simple implementation of AnonymizationKeyProvider that uses a string as the key.
type AnonymizationKeyString string
func (h AnonymizationKeyString) GetAnonymizationKey() []byte {
return []byte(h)
}
func (h AnonymizationKeyString) InitializeAnonymizationKey() error {
return nil
}
// HMACAnonymizer implements anonymization using HMAC
type HMACAnonymizer struct {
// key is the HMAC key
keyProvider AnonymizationKeyProvider
}
var _ Anonymizer = (*HMACAnonymizer)(nil)
type AnonymizationKeyProvider interface {
// InitializeAnonymizationKey initializes the anonymization key if needed.
InitializeAnonymizationKey() error
// GetHMACAnonymizerKey returns the HMAC anonymizer key.
GetAnonymizationKey() []byte
}
// NewHMACAnonymizer returns a new HMAC-based anonymizer
func NewHMACAnonymizer(keyProvider AnonymizationKeyProvider) (*HMACAnonymizer, error) {
if err := keyProvider.InitializeAnonymizationKey(); err != nil {
return nil, trace.Wrap(err, "failed to initialize anonymization key")
}
key := keyProvider.GetAnonymizationKey()
if strings.TrimSpace(string(key)) == "" {
return nil, trace.BadParameter("HMAC key must not be empty")
}
return &HMACAnonymizer{keyProvider: keyProvider}, nil
}
// Anonymize anonymizes the provided data using HMAC
func (a *HMACAnonymizer) Anonymize(data []byte) string {
k := a.keyProvider.GetAnonymizationKey()
h := hmac.New(sha256.New, k)
h.Write(data)
return base64.StdEncoding.EncodeToString(h.Sum(nil))
}
// AnonymizeString anonymizes the given string data using HMAC
func (a *HMACAnonymizer) AnonymizeString(s string) string {
return a.Anonymize([]byte(s))
}
// AnonymizeNonEmpty implements [Anonymizer].
func (a *HMACAnonymizer) AnonymizeNonEmpty(s string) []byte {
if s == "" {
return nil
}
k := a.keyProvider.GetAnonymizationKey()
h := hmac.New(sha256.New, k)
h.Write([]byte(s))
return h.Sum(nil)
}
// Teleport
// Copyright (C) 2024 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 utils
import (
"cmp"
"fmt"
"net"
"strings"
"github.com/gravitational/teleport/api/types"
scopedapp "github.com/gravitational/teleport/lib/scopes/app"
)
// AssembleAppFQDN returns the application's FQDN.
//
// If the application is running within the local cluster and it has a public
// address specified, the application's public address is used.
//
// In all other cases, i.e. if the public address is not set or the application
// is running in a remote cluster, the FQDN is formatted as
// <appName>.<localProxyDNSName>.
// A scoped app is always addressed by its computed hash label under the
// selected proxy as <hash(appName, scope)>.<localProxyDNSName>.
func AssembleAppFQDN(localClusterName string, localProxyDNSName string, appClusterName string, app types.Application) string {
if scope := app.GetScope(); scope != "" {
return scopedapp.ScopedAppPublicAddr(scope, app.GetName(), localProxyDNSName)
}
isLocalCluster := localClusterName == appClusterName
if isLocalCluster && app.GetPublicAddr() != "" && !app.GetUseAnyProxyPublicAddr() {
return app.GetPublicAddr()
}
return DefaultAppPublicAddr(app.GetName(), localProxyDNSName)
}
// DefaultAppPublicAddr returns "<appName>.<localProxyDNSName>",
// stripping a trailing port and lowercasing the host so the result
// satisfies ValidatePublicAddr.
func DefaultAppPublicAddr(appName, localProxyDNSName string) string {
if host, _, err := net.SplitHostPort(localProxyDNSName); err == nil {
localProxyDNSName = host
}
return fmt.Sprintf("%s.%s", appName, strings.ToLower(localProxyDNSName))
}
// DefaultAppFQDN returns the default routing FQDN for an app.
// proxyPublicAddrHost takes precedence; clusterName is the fallback
// when it is empty. An IP-valued proxy public_addr is used as-is,
// not replaced by clusterName.
func DefaultAppFQDN(appName, proxyPublicAddrHost, clusterName string) string {
return DefaultAppPublicAddr(appName, cmp.Or(proxyPublicAddrHost, clusterName))
}
/*
* 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 utils
import (
"archive/tar"
"bytes"
"compress/gzip"
"io/fs"
"github.com/gravitational/trace"
)
// ReadStatFS combines two interfaces: fs.ReadFileFS and fs.StatFS
// We need both when creating the archive to be able to:
// - read file contents - `ReadFile` provided by fs.ReadFileFS
// - set the correct file permissions - `Stat() ... Mode()` provided by fs.StatFS
type ReadStatFS interface {
fs.ReadFileFS
fs.StatFS
}
// CompressTarGzArchive creates a Tar Gzip archive in memory, reading the files using the provided file reader
func CompressTarGzArchive(files []string, fileReader ReadStatFS) (*bytes.Buffer, error) {
archiveBytes := &bytes.Buffer{}
gzipWriter, err := gzip.NewWriterLevel(archiveBytes, gzip.BestSpeed)
if err != nil {
return nil, trace.Wrap(err)
}
defer gzipWriter.Close()
tarWriter := tar.NewWriter(gzipWriter)
defer tarWriter.Close()
for _, filename := range files {
bs, err := fileReader.ReadFile(filename)
if err != nil {
return nil, trace.Wrap(err)
}
fileStat, err := fileReader.Stat(filename)
if err != nil {
return nil, trace.Wrap(err)
}
if err := tarWriter.WriteHeader(&tar.Header{
Name: filename,
Size: int64(len(bs)),
Mode: int64(fileStat.Mode()),
}); err != nil {
return nil, trace.Wrap(err)
}
if _, err := tarWriter.Write(bs); err != nil {
return nil, trace.Wrap(err)
}
}
return archiveBytes, nil
}
/*
* 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 aws
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"log/slog"
"net/http"
"net/textproto"
"slices"
"sort"
"strings"
"time"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/aws/arn"
v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4"
"github.com/gravitational/trace"
apievents "github.com/gravitational/teleport/api/types/events"
apiawsutils "github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/lib/utils"
)
const (
// AmazonSigV4AuthorizationPrefix is AWS Authorization prefix indicating that the request
// was signed by AWS Signature Version 4.
// https://github.com/aws/aws-sdk-go/blob/main/aws/signer/v4/v4.go#L83
// https://docs.aws.amazon.com/AmazonS3/latest/API/sigv4-auth-using-authorization-header.html
AmazonSigV4AuthorizationPrefix = "AWS4-HMAC-SHA256"
// AmzDateTimeFormat is time format used in X-Amz-Date header.
// https://github.com/aws/aws-sdk-go/blob/main/aws/signer/v4/v4.go#L84
AmzDateTimeFormat = "20060102T150405Z"
// AmzDateHeader is header name containing timestamp when signature was generated.
// https://docs.aws.amazon.com/general/latest/gr/sigv4-date-handling.html
AmzDateHeader = "X-Amz-Date"
AuthorizationHeader = "Authorization"
credentialAuthHeaderElem = "Credential"
signedHeaderAuthHeaderElem = "SignedHeaders"
signatureAuthHeaderElem = "Signature"
// AmzTargetHeader is a header containing the API target.
// Format: target_version.operation
// Example: DynamoDB_20120810.Scan
AmzTargetHeader = "X-Amz-Target"
// AmzJSON1_0 is an AWS Content-Type header that indicates the media type is JSON.
AmzJSON1_0 = "application/x-amz-json-1.0"
// AmzJSON1_1 is an AWS Content-Type header that indicates the media type is JSON.
AmzJSON1_1 = "application/x-amz-json-1.1"
// MaxRoleSessionNameLength is the maximum length of the role session name
// used by the AssumeRole call.
// https://docs.aws.amazon.com/IAM/latest/UserGuide/reference_iam-quotas.html
MaxRoleSessionNameLength = 64
// EmptyPayloadHash is the SHA-256 for an empty element (as in echo -n | sha256sum).
// PresignHTTP requires the hash of the body, but when there is no body we hash the empty string.
// https://docs.aws.amazon.com/AmazonS3/latest/API/sig-v4-header-based-auth.html
EmptyPayloadHash = "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
iamServiceName = "iam"
)
// SigV4 contains parsed content of the AWS Authorization header.
type SigV4 struct {
// KeyIS is an AWS access-key-id
KeyID string
// Date value is specified using YYYYMMDD format.
Date string
// Region is an AWS Region.
Region string
// Service is an AWS Service.
Service string
// SignedHeaders is a list of request headers that you used to compute Signature.
SignedHeaders []string
// Signature is the 256-bit Signature of the request.
Signature string
}
// ParseSigV4 AWS SigV4 credentials string sections.
// AWS SigV4 header example below adds newlines for readability only - the real
// header must be a single continuous string with commas (and optional spaces)
// between the Credential, SignedHeaders, and Signature:
// Authorization: AWS4-HMAC-SHA256 Credential=AKIAIOSFODNN7EXAMPLE/20130524/us-east-1/s3/aws4_request,
// SignedHeaders=host;range;x-amz-date,
// Signature=fe5f80f77d5fa3beca038a248ff027d0445342fe2855ddc963176630326f1024
func ParseSigV4(header string) (*SigV4, error) {
if header == "" {
return nil, trace.BadParameter("empty AWS SigV4 header")
}
if !strings.HasPrefix(header, AmazonSigV4AuthorizationPrefix+" ") {
return nil, trace.BadParameter("missing AWS SigV4 authorization algorithm")
}
header = strings.TrimPrefix(header, AmazonSigV4AuthorizationPrefix+" ")
components := strings.Split(header, ",")
if len(components) != 3 {
return nil, trace.BadParameter("expected AWS SigV4 Authorization header with 3 comma-separated components but got %d", len(components))
}
m := make(map[string]string)
for _, v := range components {
kv := strings.Split(strings.Trim(v, " "), "=")
if len(kv) != 2 {
continue
}
m[kv[0]] = kv[1]
}
authParts := strings.Split(m[credentialAuthHeaderElem], "/")
if len(authParts) != 5 {
return nil, trace.BadParameter("invalid size of %q section", credentialAuthHeaderElem)
}
signature := m[signatureAuthHeaderElem]
if signature == "" {
return nil, trace.BadParameter("invalid signature")
}
var signedHeaders []string
if v := m[signedHeaderAuthHeaderElem]; v != "" {
signedHeaders = strings.Split(v, ";")
}
return &SigV4{
KeyID: authParts[0],
Date: authParts[1],
Region: authParts[2],
Service: authParts[3],
Signature: signature,
// Split semicolon-separated list of signed headers string.
// https://docs.aws.amazon.com/AmazonS3/latest/API/sigv4-auth-using-authorization-header.html
// https://github.com/aws/aws-sdk-go/blob/main/aws/signer/v4/v4.go#L631
SignedHeaders: signedHeaders,
}, nil
}
// IsSignedByAWSSigV4 checks is the request was signed by AWS Signature Version 4 algorithm.
// https://docs.aws.amazon.com/general/latest/gr/signing_aws_api_requests.html
func IsSignedByAWSSigV4(r *http.Request) bool {
return strings.HasPrefix(r.Header.Get(AuthorizationHeader), AmazonSigV4AuthorizationPrefix)
}
// VerifyAWSSignature verifies the request signature ensuring that the request originates from tsh aws command execution
// AWS CLI signs the request with random generated credentials that are passed to LocalProxy by
// the AWSCredentials LocalProxyConfig configuration.
func VerifyAWSSignature(req *http.Request, credProvider aws.CredentialsProvider) error {
sigV4, err := ParseSigV4(req.Header.Get("Authorization"))
if err != nil {
return trace.BadParameter("%s", err)
}
// Verifies the request is signed by the expected access key ID.
credValue, err := credProvider.Retrieve(req.Context())
if err != nil {
return trace.Wrap(err)
}
if sigV4.KeyID != credValue.AccessKeyID {
return trace.AccessDenied("AccessKeyID does not match")
}
// Skip signature verification if the incoming request includes the
// "User-Agent" header when making the signature. AWS Go SDK explicitly
// skips the "User-Agent" header so it will always produce a different
// signature. Only AccessKeyID is verified above in this case.
for _, signedHeader := range sigV4.SignedHeaders {
if strings.EqualFold(signedHeader, "User-Agent") {
return nil
}
}
payloadHash, err := GetV4PayloadHash(req)
if err != nil {
return trace.Wrap(err)
}
ctx := context.Background()
reqCopy := req.Clone(ctx)
// Remove all the headers that are not present in awsCred.SignedHeaders.
filterHeaders(reqCopy, sigV4.SignedHeaders)
// Get the date that was used to create the signature of the original request
// originated from AWS CLI and reuse it as a timestamp during request signing call.
t, err := time.Parse(AmzDateTimeFormat, reqCopy.Header.Get(AmzDateHeader))
if err != nil {
return trace.BadParameter("%s", err)
}
creds, err := credProvider.Retrieve(ctx)
if err != nil {
return trace.Wrap(err)
}
// If the original request does not sign "content-length" (e.g. "aws sts
// get-caller-identity"), do not set reqCopy.ContentLength as go sdk's
// signer will forcefully sign "content-length" if it is set on the HTTP
// request.
findContentLength := slices.ContainsFunc(sigV4.SignedHeaders, func(header string) bool {
return strings.EqualFold(header, "content-length")
})
if !findContentLength {
reqCopy.ContentLength = 0
}
signer := NewSigner(sigV4.Service)
err = signer.SignHTTP(ctx, creds, reqCopy, payloadHash, sigV4.Service, sigV4.Region, t)
if err != nil {
return trace.Wrap(err)
}
localSigV4, err := ParseSigV4(reqCopy.Header.Get("Authorization"))
if err != nil {
return trace.Wrap(err)
}
// Compare the origin request AWS SigV4 signature with the signature calculated in LocalProxy based on
// AWSCredentials taken from LocalProxyConfig.
if sigV4.Signature != localSigV4.Signature {
return trace.AccessDenied("signature verification failed")
}
return nil
}
// NewSigner creates a new V4 signer.
func NewSigner(signingServiceName string) *v4.Signer {
return v4.NewSigner(func(opts *v4.SignerOptions) {
// s3 and s3control requests are signed with URL unescaped (found by
// searching "DisableURIPathEscaping" in "aws-sdk-go/service"). Both
// services use "s3" as signing name. See description of
// "DisableURIPathEscaping" for more details.
if signingServiceName == "s3" {
opts.DisableURIPathEscaping = true
}
})
}
// GetV4PayloadHash returns the V4 signing payload hash.
func GetV4PayloadHash(req *http.Request) (string, error) {
payloadHash := strings.ToUpper(req.Header.Get("x-amz-content-sha256"))
switch payloadHash {
// unsigned payload, so we use the literal content string instead of hashing
// https://docs.aws.amazon.com/AmazonS3/latest/API/sigv4-auth-using-authorization-header.html
case "UNSIGNED-PAYLOAD", "STREAMING-UNSIGNED-PAYLOAD-TRAILER":
return payloadHash, nil
default:
}
payload, err := utils.GetAndReplaceRequestBody(req)
if err != nil {
return "", trace.Wrap(err)
}
if len(payload) == 0 {
return EmptyPayloadHash, nil
}
hash := sha256.New()
hash.Write(payload)
return hex.EncodeToString(hash.Sum(nil)), nil
}
// filterHeaders removes request headers that are not in the headers list and returns the removed header keys.
func filterHeaders(r *http.Request, headers []string) []string {
keep := make(map[string]struct{})
for _, key := range headers {
keep[textproto.CanonicalMIMEHeaderKey(key)] = struct{}{}
}
var removed []string
out := make(http.Header)
for key, vals := range r.Header {
if _, ok := keep[textproto.CanonicalMIMEHeaderKey(key)]; ok {
out[key] = vals
continue
}
removed = append(removed, key)
}
r.Header = out
return removed
}
// FilterAWSRoles returns role ARNs from the provided list that belong to the
// specified AWS account ID.
//
// If AWS account ID is empty, all valid AWS IAM roles are returned.
func FilterAWSRoles(arns []string, accountID string) (result Roles) {
for _, roleARN := range arns {
parsed, err := ParseRoleARN(roleARN)
if err != nil {
slog.WarnContext(context.Background(), "Skipping invalid AWS role ARN.", "error", err)
continue
}
if accountID != "" && parsed.AccountID != accountID {
continue
}
// In AWS convention, the display of the role is the last
// /-delineated substring.
//
// Example ARNs:
// arn:aws:iam::1234567890:role/EC2FullAccess (display: EC2FullAccess)
// arn:aws:iam::1234567890:role/path/to/customrole (display: customrole)
parts := strings.Split(parsed.Resource, "/")
result = append(result, Role{
Name: strings.Join(parts[1:], "/"),
Display: parts[len(parts)-1],
ARN: roleARN,
AccountID: parsed.AccountID,
})
}
return result
}
// Role describes an AWS IAM role for AWS console access.
type Role struct {
// Name is the full role name with the entire path.
Name string `yaml:"name" json:"name"`
// Display is the role display name.
Display string `yaml:"display" json:"display"`
// ARN is the full role ARN.
ARN string `yaml:"arn" json:"arn"`
// AccountID is the AWS Account ID this role refers to.
AccountID string `yaml:"accountId" json:"accountId"`
// RequiresRequest indicates whether this role requires an access request
// to be used.
RequiresRequest bool `yaml:"requiresRequest,omitempty" json:"requiresRequest,omitempty"`
}
// Roles is a slice of roles.
type Roles []Role
// Sort sorts the roles by their display names.
func (roles Roles) Sort() {
sort.SliceStable(roles, func(x, y int) bool {
return strings.ToLower(roles[x].Display) < strings.ToLower(roles[y].Display)
})
}
// FindRoleByARN finds the role with the provided ARN.
func (roles Roles) FindRoleByARN(arn string) (Role, bool) {
for _, role := range roles {
if role.ARN == arn {
return role, true
}
}
return Role{}, false
}
// FindRolesByName finds all roles matching the provided name.
func (roles Roles) FindRolesByName(name string) (result Roles) {
for _, role := range roles {
// Match either full name or the display name.
if role.Display == name || role.Name == name {
result = append(result, role)
}
}
return
}
// UnmarshalRequestBody reads and unmarshals a JSON request body into a protobuf Struct wrapper.
// If the request is not a recognized AWS JSON media type, or the body cannot be read, or the body
// is not valid JSON, then this function returns a nil value and an error.
// The protobuf Struct wrapper is useful for serializing JSON into a protobuf, because otherwise when the
// protobuf is marshaled it will re-marshall a JSON string field with escape characters or base64 encode
// a []byte field.
// Examples showing differences:
// - JSON string in proto: `{"Table": "some-table"}` --marshal to JSON--> `"{\"Table\": \"some-table\"}"`
// - bytes in proto: []byte --marshal to JSON--> `eyJUYWJsZSI6ICJzb21lLXRhYmxlIn0K` (base64 encoded)
// - *Struct in proto: *Struct --marshal to JSON--> `{"Table": "some-table"}` (unescaped JSON)
func UnmarshalRequestBody(req *http.Request) (*apievents.Struct, error) {
contentType := req.Header.Get("Content-Type")
if !isJSON(contentType) {
return nil, trace.BadParameter("invalid JSON request Content-Type: %q", contentType)
}
jsonBody, err := utils.GetAndReplaceRequestBody(req)
if err != nil {
return nil, trace.Wrap(err)
}
s := &apievents.Struct{}
if err := s.UnmarshalJSON(jsonBody); err != nil {
return nil, trace.Wrap(err)
}
return s, nil
}
// isJSON returns true if the Content-Type is recognized as standard JSON or any non-standard
// Amazon Content-Type header that indicates JSON media type.
func isJSON(contentType string) bool {
switch contentType {
case "application/json", AmzJSON1_0, AmzJSON1_1:
return true
default:
return false
}
}
// BuildRoleARN constructs a string AWS ARN from a username, region, and account ID.
// If username is an AWS ARN, this function checks that the ARN is an AWS IAM Role ARN
// in the correct partition and account.
func BuildRoleARN(username, region, accountID string) (string, error) {
partition := apiawsutils.GetPartitionFromRegion(region)
if arn.IsARN(username) {
// sanity check the given username role ARN.
parsed, err := ParseRoleARN(username)
if err != nil {
return "", trace.Wrap(err)
}
// don't check for empty accountID - callers do not always pass an account ID,
// and it's only absolutely required if we need to build the role ARN below.
if err := CheckARNPartitionAndAccount(parsed, partition, accountID); err != nil {
return "", trace.Wrap(err)
}
return username, nil
}
resource := username
if !IsPartialRoleARN(resource) {
resource = fmt.Sprintf("role/%s", username)
}
roleARN := arn.ARN{
Partition: partition,
Service: iamServiceName,
AccountID: accountID,
Resource: resource,
}
if err := apiawsutils.CheckRoleARN(roleARN.String()); err != nil {
return "", trace.Wrap(err)
}
return roleARN.String(), nil
}
// ValidateRoleARNAndExtractRoleName validates the role ARN and extracts the
// short role name from it.
func ValidateRoleARNAndExtractRoleName(roleARN, wantPartition, wantAccountID string) (string, error) {
role, err := ParseRoleARN(roleARN)
if err != nil {
return "", trace.Wrap(err)
}
if err := CheckARNPartitionAndAccount(role, wantPartition, wantAccountID); err != nil {
return "", trace.Wrap(err)
}
return strings.TrimPrefix(role.Resource, "role/"), nil
}
// ParseRoleARN parses an AWS ARN and checks that the ARN is
// for an IAM Role resource.
func ParseRoleARN(roleARN string) (*arn.ARN, error) {
role, err := arn.Parse(roleARN)
if err != nil {
return nil, trace.BadParameter("invalid AWS ARN: %v", err)
}
if err := checkRoleARN(&role); err != nil {
return nil, trace.Wrap(err)
}
return &role, nil
}
// checkRoleARN returns whether a parsed ARN is for an IAM Role resource.
// Example role ARN: arn:aws:iam::123456789012:role/some-role-name
func checkRoleARN(parsed *arn.ARN) error {
parts := strings.Split(parsed.Resource, "/")
if parts[0] != "role" || parsed.Service != iamServiceName {
return trace.BadParameter("%q is not an AWS IAM role ARN", parsed)
}
if len(parts) < 2 || len(parts[len(parts)-1]) == 0 {
return trace.BadParameter("%q is missing AWS IAM role name", parsed)
}
if err := apiawsutils.IsValidAccountID(parsed.AccountID); err != nil {
return trace.BadParameter("%q invalid account ID: %v", parsed, err)
}
return nil
}
// CheckARNPartitionAndAccount checks an AWS ARN against an expected AWS partition and account ID.
// An empty expected AWS partition or account ID is not checked.
func CheckARNPartitionAndAccount(ARN *arn.ARN, wantPartition, wantAccountID string) error {
if ARN.Partition != wantPartition && wantPartition != "" {
return trace.BadParameter("expected AWS partition %q but got %q", wantPartition, ARN.Partition)
}
if ARN.AccountID != wantAccountID && wantAccountID != "" {
return trace.BadParameter("expected AWS account ID %q but got %q", wantAccountID, ARN.AccountID)
}
return nil
}
// IsRoleARN returns true if the provided string is a AWS role ARN.
func IsRoleARN(roleARN string) bool {
if _, err := ParseRoleARN(roleARN); err == nil {
return true
}
return IsPartialRoleARN(roleARN)
}
// IsPartialRoleARN returns true if the provided role ARN only contains the
// resource name.
func IsPartialRoleARN(roleARN string) bool {
return strings.HasPrefix(roleARN, "role/")
}
// IsUserARN returns true if the provided string is a AWS user ARN.
func IsUserARN(userARN string) bool {
resourceName := userARN
if parsed, err := arn.Parse(userARN); err == nil {
resourceName = parsed.Resource
}
return strings.HasPrefix(resourceName, "user/")
}
// PolicyARN returns the ARN representation of an AWS IAM Policy.
func PolicyARN(partition, accountID, policy string) string {
return iamResourceARN(partition, accountID, "policy", policy)
}
// RoleARN returns the ARN representation of an AWS IAM Role.
func RoleARN(partition, accountID, role string) string {
return iamResourceARN(partition, accountID, "role", role)
}
func iamResourceARN(partition, accountID, resourceType, resourceName string) string {
return arn.ARN{
Partition: partition,
Service: "iam",
AccountID: accountID,
Resource: fmt.Sprintf("%s/%s", resourceType, resourceName),
}.String()
}
/*
* 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 aws
import (
"context"
"github.com/aws/aws-sdk-go-v2/feature/ec2/imds"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
config "github.com/gravitational/teleport/lib/cloud/aws/config"
"github.com/gravitational/teleport/lib/utils"
)
// GetRawEC2IdentityDocument fetches the PKCS7 RSA2048 InstanceIdentityDocument
// from the IMDS for this EC2 instance.
func GetRawEC2IdentityDocument(ctx context.Context) ([]byte, error) {
cfg, err := config.LoadDefaultConfig(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
imdsClient := imds.NewFromConfig(cfg)
output, err := imdsClient.GetDynamicData(ctx, &imds.GetDynamicDataInput{
Path: "instance-identity/rsa2048",
})
if err != nil {
return nil, trace.Wrap(err)
}
iidBytes, err := utils.ReadAtMost(output.Content, teleport.MaxHTTPResponseSize)
if err != nil {
return nil, trace.Wrap(err)
}
if err := output.Content.Close(); err != nil {
return nil, trace.Wrap(err)
}
return iidBytes, nil
}
func GetEC2InstanceIdentityDocument(ctx context.Context) (*imds.InstanceIdentityDocument, error) {
// fetch the raw IID
cfg, err := config.LoadDefaultConfig(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
imdsClient := imds.NewFromConfig(cfg)
output, err := imdsClient.GetInstanceIdentityDocument(ctx, nil)
if err != nil {
return nil, trace.Wrap(err)
}
return &output.InstanceIdentityDocument, nil
}
// GetEC2NodeID returns the node ID to use for this EC2 instance when using
// Simplified Node Joining.
func GetEC2NodeID(ctx context.Context) (string, error) {
// fetch the raw IID
iid, err := GetEC2InstanceIdentityDocument(ctx)
if err != nil {
return "", trace.Wrap(err)
}
return NodeIDFromIID(iid), nil
}
// NodeIDFromIID returns the node ID that must be used for nodes joining with
// the given Instance Identity Document.
func NodeIDFromIID(iid *imds.InstanceIdentityDocument) string {
return iid.AccountID + "-" + iid.InstanceID
}
/*
* 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 aws
import (
"context"
"errors"
"io"
"strings"
awsv2 "github.com/aws/aws-sdk-go-v2/aws"
"github.com/aws/aws-sdk-go-v2/feature/s3/transfermanager"
s3types "github.com/aws/aws-sdk-go-v2/service/s3/types"
"github.com/aws/smithy-go"
"github.com/gravitational/trace"
)
// ConvertS3Error wraps S3 error and returns trace equivalent
// It works on both sdk v1 and v2.
func ConvertS3Error(err error) error {
if err == nil {
return nil
}
var noSuchKey *s3types.NoSuchKey
if errors.As(err, &noSuchKey) {
return trace.NotFound("%s", noSuchKey)
}
var noSuchBucket *s3types.NoSuchBucket
if errors.As(err, &noSuchBucket) {
return trace.NotFound("%s", noSuchBucket)
}
var noSuchUpload *s3types.NoSuchUpload
if errors.As(err, &noSuchUpload) {
return trace.NotFound("%s", noSuchUpload)
}
var bucketAlreadyExists *s3types.BucketAlreadyExists
if errors.As(err, &bucketAlreadyExists) {
return trace.AlreadyExists("%s", bucketAlreadyExists.Error())
}
var bucketAlreadyOwned *s3types.BucketAlreadyOwnedByYou
if errors.As(err, &bucketAlreadyOwned) {
return trace.AlreadyExists("%s", bucketAlreadyOwned.Error())
}
var notFound *s3types.NotFound
if errors.As(err, ¬Found) {
return trace.NotFound("%s", notFound)
}
var opError *smithy.OperationError
if errors.As(err, &opError) && strings.Contains(opError.Err.Error(), "FIPS") {
return trace.BadParameter("%s", opError)
}
return err
}
// s3V2FileWriter can be used to upload data to s3 via io.WriteCloser interface.
type s3V2FileWriter struct {
// uploadFinisherErrChan is used to wait for completed upload as well as
// sending error message.
uploadFinisherErrChan <-chan error
pipeWriter *io.PipeWriter
pipeReader *io.PipeReader
}
// NewS3V2FileWriter created s3V2FileWriter. Close method on writer should be called
// to make sure that reader has finished.
func NewS3V2FileWriter(ctx context.Context, s3Client transfermanager.S3APIClient, bucket, key string, uploaderOptions []func(*transfermanager.Options), putObjectInputOptions ...func(*transfermanager.UploadObjectInput)) (*s3V2FileWriter, error) {
client := transfermanager.New(s3Client, uploaderOptions...)
pr, pw := io.Pipe()
uploadParams := &transfermanager.UploadObjectInput{
Bucket: awsv2.String(bucket),
Key: awsv2.String(key),
Body: pr,
}
for _, f := range putObjectInputOptions {
f(uploadParams)
}
uploadFinisherErrChan := make(chan error)
go func() {
defer close(uploadFinisherErrChan)
_, err := client.UploadObject(ctx, uploadParams)
if err != nil {
pr.CloseWithError(err)
}
uploadFinisherErrChan <- trace.Wrap(err)
}()
return &s3V2FileWriter{
uploadFinisherErrChan: uploadFinisherErrChan,
pipeWriter: pw,
pipeReader: pr,
}, nil
}
// Write bytes from in to the connected pipe.
func (s *s3V2FileWriter) Write(in []byte) (int, error) {
bytesWritten, writeError := s.pipeWriter.Write(in)
if writeError != nil {
s.pipeWriter.CloseWithError(writeError)
return bytesWritten, writeError
}
return bytesWritten, nil
}
// Close signals write completion and cleans up any
// open streams. Will block until pending uploads are complete.
func (s *s3V2FileWriter) Close() error {
wCloseErr := s.pipeWriter.Close()
// wait for reader to finish, it will be triggered by writer.Close
readerErr := <-s.uploadFinisherErrChan
rCloseErr := s.pipeReader.Close()
return trace.Wrap(trace.NewAggregate(wCloseErr, readerErr, rCloseErr))
}
// CreateBucketConfiguration creates the default CreateBucketConfiguration.
func CreateBucketConfiguration(region string) *s3types.CreateBucketConfiguration {
// No location constraint wanted for us-east-1 because it is the default and
// AWS has decided, in all their infinite wisdom, that the CreateBucket API
// should fail if you explicitly pass the default location constraint.
if region == "us-east-1" {
return nil
}
return &s3types.CreateBucketConfiguration{
LocationConstraint: s3types.BucketLocationConstraint(region),
}
}
/*
* 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 aws
import (
"context"
"io"
"net/http"
"time"
"github.com/aws/aws-sdk-go-v2/aws"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/lib/utils"
)
// SigningCtx contains AWS SigV4 signing context parameters.
type SigningCtx struct {
// Clock is used to override time in tests.
Clock clockwork.Clock
// Credentials provides AWS credentials.
Credentials aws.CredentialsProvider
// SigningName is the AWS signing service name.
SigningName string
// SigningRegion is the AWS region to sign a request for.
SigningRegion string
}
// Check checks signing context parameters.
func (sc *SigningCtx) Check() error {
switch {
case sc.Credentials == nil:
return trace.BadParameter("missing AWS credentials")
case sc.SigningName == "":
return trace.BadParameter("missing AWS signing name")
case sc.SigningRegion == "":
return trace.BadParameter("missing AWS signing region")
}
return nil
}
// SignRequest creates a new HTTP request and rewrites the header from the original request and returns a new
// HTTP request signed by STS AWS API.
// Signing steps:
// 1) Decode Authorization Header. Authorization Header example:
//
// Authorization: AWS4-HMAC-SHA256
// Credential=AKIAIOSFODNN7EXAMPLE/20130524/us-east-1/s3/aws4_request,
// SignedHeaders=host;range;x-amz-date,
// Signature=fe5f80f77d5fa3beca038a248ff027d0445342fe2855ddc963176630326f1024
//
// 2. Extract credential section from credential Authorization Header.
// 3. Extract aws-region and aws-service from the credential section.
// 4. Build AWS API endpoint based on extracted aws-region and aws-service fields.
// Not that for endpoint resolving the https://github.com/aws/aws-sdk-go/aws/endpoints/endpoints.go
// package is used and when Amazon releases a new API the dependency update is needed.
// 5. Sign HTTP request.
func SignRequest(ctx context.Context, req *http.Request, signCtx *SigningCtx) (*http.Request, error) {
if signCtx == nil {
return nil, trace.BadParameter("missing signing context")
}
if err := signCtx.Check(); err != nil {
return nil, trace.Wrap(err)
}
payloadHash, err := GetV4PayloadHash(req)
if err != nil {
return nil, trace.Wrap(err)
}
reqCopy := req.Clone(ctx)
reqCopy.Body = io.NopCloser(req.Body)
// Only keep the headers signed in the original request for signing. This
// not only avoids signing extra headers injected by Teleport along the
// way, but also preserves the signing logic of the original AWS client.
//
// For example, Athena ODBC driver sends query requests with "Expect:
// 100-continue" headers without being signed, otherwise the Athena service
// would reject the requests.
unsignedHeaders := removeUnsignedHeaders(reqCopy)
creds, err := signCtx.Credentials.Retrieve(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
signer := NewSigner(signCtx.SigningName)
err = signer.SignHTTP(ctx, creds, reqCopy, payloadHash, signCtx.SigningName, signCtx.SigningRegion, time.Now())
if err != nil {
return nil, trace.Wrap(err)
}
// copy removed headers back to the request after signing it, but don't copy the old Authorization header.
copyHeaders(reqCopy, req, utils.RemoveFromSlice(unsignedHeaders, "Authorization"))
return reqCopy, nil
}
// removeUnsignedHeaders removes and returns header keys that are not included in SigV4 SignedHeaders.
// If the request is not already signed, then no headers are removed.
func removeUnsignedHeaders(reqCopy *http.Request) []string {
// check if the request is already signed.
authHeader := reqCopy.Header.Get("Authorization")
sig, err := ParseSigV4(authHeader)
if err != nil {
return nil
}
return filterHeaders(reqCopy, sig.SignedHeaders)
}
// copyHeaders copies headers from src request to dst request, using a list of header keys to copy.
func copyHeaders(dst *http.Request, src *http.Request, keys []string) {
for _, k := range keys {
if vals, ok := src.Header[k]; ok {
dst.Header[k] = vals
}
}
}
/*
* 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 utils
import (
"golang.org/x/crypto/bcrypt"
)
const maxInputSize = 72
// truncateToMaxSize Make sure input is truncated to the maximum length crypto accepts. Crypto changed the behavior
// from ignoring the extra input to returning an error, this truncation is necessary to maintain compatibility with
// customers who have long passwords, or more commonly our recovery codes.
func truncateToMaxSize(input []byte) []byte {
if len(input) > maxInputSize {
return input[:maxInputSize]
}
return input
}
// BcryptFromPassword delegates to bcrypt.GenerateFromPassword, but maintains the prior behavior of only hashing the
// first 72 bytes. BCrypt as an algorithm can not hash inputs > 72 bytes.
func BcryptFromPassword(password []byte, cost int) ([]byte, error) {
return bcrypt.GenerateFromPassword(truncateToMaxSize(password), cost)
}
/*
* 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 utils
import (
"sync"
)
// NewCloseBroadcaster returns new instance of close broadcaster
func NewCloseBroadcaster() *CloseBroadcaster {
return &CloseBroadcaster{
C: make(chan struct{}),
}
}
// CloseBroadcaster is a helper struct
// that implements io.Closer and uses channel
// to broadcast its closed state once called
type CloseBroadcaster struct {
sync.Once
C chan struct{}
}
// Close closes channel (once) to start broadcasting its closed state
func (b *CloseBroadcaster) Close() error {
b.Do(func() {
close(b.C)
})
return nil
}
/*
Copyright 2016 SPIFFE Authors
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.
*/
package utils
import (
"bytes"
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/binary"
"encoding/pem"
"fmt"
"math/big"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/api/utils/tlsutils"
"github.com/gravitational/teleport/lib/cryptosuites"
)
// ParseKeyStorePEM parses signing key store from PEM encoded key pair
func ParseKeyStorePEM(keyPEM, certPEM string) (*KeyStore, error) {
_, err := tlsutils.ParseCertificatePEM([]byte(certPEM))
if err != nil {
return nil, trace.Wrap(err)
}
key, err := keys.ParsePrivateKey([]byte(keyPEM))
if err != nil {
return nil, trace.Wrap(err)
}
rsaKey, ok := key.Signer.(*rsa.PrivateKey)
if !ok {
return nil, trace.BadParameter("key of type %T is not supported, only RSA keys are supported", key)
}
certASN, _ := pem.Decode([]byte(certPEM))
if certASN == nil {
return nil, trace.BadParameter("expected PEM-encoded block")
}
return &KeyStore{privateKey: rsaKey, cert: certASN.Bytes}, nil
}
// KeyStore is used to sign and decrypt data using X509 digital signatures.
type KeyStore struct {
privateKey *rsa.PrivateKey
cert []byte
}
// GetKeyPair implements goxmldsig.X509KeyPair.
func (ks *KeyStore) GetKeyPair() (*rsa.PrivateKey, []byte, error) {
return ks.privateKey, ks.cert, nil
}
// GenerateSelfSignedSigningCert is an alias of GenerateRSASelfSignedSigningCert
// used due to references in teleport.e.
var GenerateSelfSignedSigningCert = GenerateRSASelfSignedSigningCert
// GenerateRSASelfSignedSigningCert generates an RSA self-signed certificate used
// for digital signatures.
// This should only be used for the SAML implementation and tests.
func GenerateRSASelfSignedSigningCert(entity pkix.Name, dnsNames []string, ttl time.Duration) ([]byte, []byte, error) {
priv, err := cryptosuites.GenerateKeyWithAlgorithm(cryptosuites.RSA2048)
if err != nil {
return nil, nil, trace.Wrap(err)
}
rsaPriv := priv.(*rsa.PrivateKey)
// to account for clock skew
notBefore := time.Now().Add(-2 * time.Minute)
notAfter := notBefore.Add(ttl)
serialNumberLimit := new(big.Int).Lsh(big.NewInt(1), 128)
serialNumber, err := rand.Int(rand.Reader, serialNumberLimit)
if err != nil {
return nil, nil, trace.Wrap(err)
}
template := x509.Certificate{
SerialNumber: serialNumber,
Issuer: entity,
Subject: entity,
NotBefore: notBefore,
NotAfter: notAfter,
KeyUsage: x509.KeyUsageDigitalSignature,
BasicConstraintsValid: true,
DNSNames: dnsNames,
}
derBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &rsaPriv.PublicKey, rsaPriv)
if err != nil {
return nil, nil, trace.Wrap(err)
}
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(rsaPriv)})
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: derBytes})
return keyPEM, certPEM, nil
}
// ParsePrivateKeyPEM parses PEM-encoded private key.
// Prefer [keys.ParsePrivateKey], this will be deleted after references are removed from teleport.e.
func ParsePrivateKeyPEM(bytes []byte) (crypto.Signer, error) {
return keys.ParsePrivateKey(bytes)
}
// VerifyCertificateExpiryWithLeeway checks the certificate's expiration status
// with leeway. The provided leeway value is added to the current time and can
// be used to account for potential client-side clock drift. Clients validating
// a certificate is valid to attempt to connect to a server should use positive
// leeway; servers that might want to give clients leeway should use a negative
// value.
func VerifyCertificateExpiryWithLeeway(c *x509.Certificate, clock clockwork.Clock, leeway time.Duration) error {
if clock == nil {
clock = clockwork.NewRealClock()
}
now := clock.Now().Add(leeway)
if now.Before(c.NotBefore) {
return x509.CertificateInvalidError{
Cert: c,
Reason: x509.Expired,
Detail: fmt.Sprintf("current time %s is before %s", now.UTC().Format(time.RFC3339), c.NotBefore.UTC().Format(time.RFC3339)),
}
}
if now.After(c.NotAfter) {
return x509.CertificateInvalidError{
Cert: c,
Reason: x509.Expired,
Detail: fmt.Sprintf("current time %s is after %s", now.UTC().Format(time.RFC3339), c.NotAfter.UTC().Format(time.RFC3339)),
}
}
return nil
}
// VerifyCertificateExpiryWithLeeway checks the certificate's expiration status
// with zero leeway.
func VerifyCertificateExpiry(c *x509.Certificate, clock clockwork.Clock) error {
return trace.Wrap(VerifyCertificateExpiryWithLeeway(c, clock, 0))
}
// VerifyTLSCertLeafExpiry checks a TLS certificate's expiration status.
func VerifyTLSCertLeafExpiry(cert tls.Certificate, clock clockwork.Clock) error {
leaf, err := TLSCertLeaf(cert)
if err != nil {
return trace.Wrap(err)
}
return trace.Wrap(VerifyCertificateExpiry(leaf, clock))
}
// VerifyCertificateChain reads in chain of certificates and makes sure the
// chain from leaf to root is valid. This ensures that clients (web browsers
// and CLI) won't have problem validating the chain.
func VerifyCertificateChain(certificateChain []*x509.Certificate) error {
// chain needs at least one certificate
if len(certificateChain) == 0 {
return trace.BadParameter("need at least one certificate in chain")
}
// extract leaf of certificate chain. it is safe to index into the chain here
// because readCertificateChain always returns a valid chain with at least
// one certificate.
leaf := certificateChain[0]
// extract intermediate certificate chain.
intermediates := x509.NewCertPool()
if len(certificateChain) > 1 {
for _, v := range certificateChain[1:] {
intermediates.AddCert(v)
}
}
// verify certificate chain, roots is nil which will cause us to to use the
// system roots.
opts := x509.VerifyOptions{
Intermediates: intermediates,
}
_, err := leaf.Verify(opts)
if err != nil {
return trace.Wrap(err)
}
return nil
}
// IsSelfSigned checks if the certificate is a self-signed certificate. To check
// if a certificate is self-signed, we make sure that only one certificate is in
// the chain and that its Subject and Issuer match.
func IsSelfSigned(certificateChain []*x509.Certificate) bool {
if len(certificateChain) != 1 {
return false
}
return bytes.Equal(certificateChain[0].RawSubject, certificateChain[0].RawIssuer)
}
// ReadCertificates parses PEM encoded bytes that can contain one or
// multiple certificates and returns a slice of x509.Certificate.
func ReadCertificates(certificateChainBytes []byte) ([]*x509.Certificate, error) {
var (
certificateBlock *pem.Block
certificates [][]byte
)
remainingBytes := bytes.TrimSpace(certificateChainBytes)
for {
certificateBlock, remainingBytes = pem.Decode(remainingBytes)
if certificateBlock == nil {
return nil, trace.NotFound("no PEM data found")
}
if t := certificateBlock.Type; t != pemBlockCertificate {
return nil, trace.BadParameter("expecting certificate, but found %v", t)
}
certificates = append(certificates, certificateBlock.Bytes)
if len(remainingBytes) == 0 {
break
}
}
// build concatenated certificates into a buffer
var buf bytes.Buffer
for _, cc := range certificates {
_, err := buf.Write(cc)
if err != nil {
return nil, trace.Wrap(err)
}
}
// parse the buffer and get a slice of x509.Certificates.
x509Certs, err := x509.ParseCertificates(buf.Bytes())
if err != nil {
return nil, trace.Wrap(err)
}
return x509Certs, nil
}
// ReadCertificatesFromPath parses PEM encoded certificates from provided path.
func ReadCertificatesFromPath(path string) ([]*x509.Certificate, error) {
bytes, err := ReadPath(path)
if err != nil {
return nil, trace.Wrap(err)
}
certs, err := ReadCertificates(bytes)
if err != nil {
return nil, trace.Wrap(err)
}
return certs, nil
}
// NewCertPoolFromPath creates a new x509.CertPool from provided path.
func NewCertPoolFromPath(path string) (*x509.CertPool, error) {
// x509.CertPool.AppendCertsFromPEM skips parse errors. Using our own
// implementation here to be more strict.
cas, err := ReadCertificatesFromPath(path)
if err != nil {
return nil, trace.Wrap(err)
}
pool := x509.NewCertPool()
for _, ca := range cas {
pool.AddCert(ca)
}
return pool, nil
}
// TLSCertLeaf is a helper function that extracts the parsed leaf *x509.Certificate
// from a tls.Certificate.
// If the leaf certificate is not parsed already, then this function parses it.
func TLSCertLeaf(cert tls.Certificate) (*x509.Certificate, error) {
if cert.Leaf != nil {
return cert.Leaf, nil
}
if len(cert.Certificate) < 1 {
return nil, trace.NotFound("invalid certificate length")
}
x509cert, err := x509.ParseCertificate(cert.Certificate[0])
return x509cert, trace.Wrap(err)
}
// InitCertLeaf initializes the Leaf field for each cert in a slice of certs,
// to reduce per-handshake processing.
// Typically, servers should avoid doing this since it will
// consume more memory.
func InitCertLeaf(cert *tls.Certificate) error {
leaf, err := TLSCertLeaf(*cert)
if err != nil {
return trace.Wrap(err)
}
cert.Leaf = leaf
return nil
}
const pemBlockCertificate = "CERTIFICATE"
// CreateCertificateBLOB creates Certificate BLOB
// It has following structure:
//
// CertificateBlob {
// PropertyID: u32, little endian,
// Reserved: u32, little endian, must be set to 0x01 0x00 0x00 0x00
// Length: u32, little endian
// Value: certificate data
// }
//
// Documentation on this structure is a little thin, but one with the structure
// exists in [MS-GPEF]. This doesn't list the `PropertyID` we use below, however
// some references can be found scattered about the internet such as [here].
//
// [MS-GPEF]: https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-gpef/e051aba9-c9df-4f82-a42a-c13012c9d381
// [here]: https://github.com/diyinfosec/010-Editor/blob/master/WINDOWS_CERTIFICATE_BLOB.bt
func CreateCertificateBLOB(certData []byte) []byte {
buf := new(bytes.Buffer)
buf.Grow(len(certData) + 12)
// PropertyID for certificate is 32
binary.Write(buf, binary.LittleEndian, int32(32))
binary.Write(buf, binary.LittleEndian, int32(1))
binary.Write(buf, binary.LittleEndian, int32(len(certData)))
buf.Write(certData)
return buf.Bytes()
}
/*
* 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 utils
import (
"sync"
"github.com/gravitational/trace"
)
// CircularBuffer implements an in-memory circular buffer of predefined size
type CircularBuffer struct {
sync.Mutex
buf []float64
start int
end int
size int
}
// NewCircularBuffer returns a new instance of a circular buffer that will hold
// size elements before it rotates
func NewCircularBuffer(size int) (*CircularBuffer, error) {
if size <= 0 {
return nil, trace.BadParameter("circular buffer size should be > 0")
}
buf := &CircularBuffer{
buf: make([]float64, size),
start: -1,
end: -1,
size: 0,
}
return buf, nil
}
// Data returns the most recent n elements in the correct order
func (t *CircularBuffer) Data(n int) []float64 {
t.Lock()
defer t.Unlock()
if n <= 0 || t.size == 0 {
return nil
}
// skip first N items so that the most recent are always provided
start := t.start
if n < t.size {
start = (t.start + (t.size - n)) % len(t.buf)
}
if start <= t.end {
return t.buf[start : t.end+1]
}
return append(t.buf[start:], t.buf[:t.end+1]...)
}
// Add pushes a new item onto the buffer
func (t *CircularBuffer) Add(d float64) {
t.Lock()
defer t.Unlock()
if t.size == 0 {
t.start = 0
t.end = 0
t.size = 1
} else if t.size < len(t.buf) {
t.end++
t.size++
} else {
t.end = t.start
t.start = (t.start + 1) % len(t.buf)
}
t.buf[t.end] = d
}
/*
* 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 utils
import (
"bytes"
"context"
"crypto/x509"
"errors"
"fmt"
"io"
"log/slog"
"os"
"runtime"
"strconv"
"strings"
"unicode"
"github.com/alecthomas/kingpin/v2"
"github.com/gravitational/trace"
"golang.org/x/term"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// LoggingPurpose specifies which kind of application logging is
// to be configured for.
type LoggingPurpose int
const (
// LoggingForDaemon configures logging for non-user interactive applications (teleport, tbot, tsh deamon).
LoggingForDaemon LoggingPurpose = iota
// LoggingForCLI configures logging for user face utilities (tctl, tsh).
LoggingForCLI
// LoggingForMCP configures logging for MCP servers.
LoggingForMCP
)
// LoggingFormat defines the possible logging output formats.
type LoggingFormat = string
const (
// LogFormatJSON configures logs to be emitted in json.
LogFormatJSON LoggingFormat = "json"
// LogFormatText configures logs to be emitted in a human readable text format.
LogFormatText LoggingFormat = "text"
)
type logOpts struct {
format LoggingFormat
// osLogSubsystem is the subsystem used for all loggers created by this process
// when sending logs to os_log on macOS. If empty, os_log won't be used.
osLogSubsystem string
}
// LoggerOption enables customizing the global logger.
type LoggerOption func(opts *logOpts)
// WithLogFormat initializes the default logger with the provided format.
func WithLogFormat(format LoggingFormat) LoggerOption {
return func(opts *logOpts) {
opts.format = format
}
}
func WithOSLog(subsystem string) LoggerOption {
return func(opts *logOpts) {
opts.osLogSubsystem = subsystem
}
}
// IsTerminal checks whether writer is a terminal
func IsTerminal(w io.Writer) bool {
switch v := w.(type) {
case *os.File:
return term.IsTerminal(int(v.Fd()))
default:
return false
}
}
// InitLogger configures the global logger for a given purpose / verbosity level
func InitLogger(purpose LoggingPurpose, level slog.Level, opts ...LoggerOption) (*slog.Logger, error) {
var o logOpts
for _, opt := range opts {
opt(&o)
}
// If debug or trace logging is not enabled for CLIs,
// then discard all log output.
if purpose == LoggingForCLI && level > slog.LevelDebug {
logger := slog.New(slog.DiscardHandler)
slog.SetDefault(logger)
return logger, nil
}
var output string
switch {
case o.osLogSubsystem != "":
output = logutils.LogOutputOSLog
case purpose == LoggingForMCP:
output = logutils.LogOutputMCP
o.format = LogFormatJSON
}
logger, _, _, err := logutils.Initialize(logutils.Config{
Severity: level.String(),
Format: o.format,
EnableColors: IsTerminal(os.Stderr),
Output: output,
OSLogSubsystem: o.osLogSubsystem,
})
return logger, trace.Wrap(err)
}
// FatalError is for CLI front-ends: it detects gravitational/trace debugging
// information, sends it to the logger, strips it off and prints a clean message to stderr
func FatalError(err error) {
fmt.Fprint(os.Stderr, UserMessageFromError(err))
os.Exit(1)
}
// GetIterations provides a simple way to add iterations to the test
// by setting environment variable "ITERATIONS", by default it returns 1
func GetIterations() int {
out := os.Getenv(teleport.IterationsEnvVar)
if out == "" {
return 1
}
iter, err := strconv.Atoi(out)
if err != nil {
panic(err)
}
slog.DebugContext(context.Background(), "Running tests multiple times due to presence of ITERATIONS environment variable", "iterations", iter)
return iter
}
// UserMessageFromError returns user-friendly error message from error.
// The error message will be formatted for output depending on the debug
// flag and will always end with a new line.
func UserMessageFromError(err error) string {
if err == nil {
return ""
}
if slog.Default().Enabled(context.Background(), slog.LevelDebug) {
msg := trace.DebugReport(err)
if !strings.HasSuffix(msg, "\n") {
msg += "\n"
}
return msg
}
var buf bytes.Buffer
if runtime.GOOS == constants.WindowsOS {
// TODO(timothyb89): Due to complications with globally enabling +
// properly resetting Windows terminal ANSI processing, for now we just
// disable color output. Otherwise, raw ANSI escapes will be visible to
// users.
fmt.Fprint(&buf, "ERROR: ")
} else {
fmt.Fprint(&buf, Color(Red, "ERROR: "))
}
formatErrorWriter(err, &buf)
return buf.String()
}
// FormatErrorWithNewline returns user friendly error message from error.
// The message will always end with a new line.
func FormatErrorWithNewline(err error) string {
var buf bytes.Buffer
formatErrorWriter(err, &buf)
return buf.String()
}
// formatErrorWriter formats the specified error into the provided writer.
// The error message is escaped if necessary. A newline is added if the
// error text does not end with a newline.
func formatErrorWriter(err error, w io.Writer) {
if err == nil {
return
}
msg := trace.UserMessage(err)
if certErr := formatCertError(err); certErr != "" {
msg = certErr
}
// Error can be of type trace.proxyError where error message didn't get captured.
if msg == "" {
fmt.Fprintln(w, "please check Teleport's log for more details")
return
}
msg = AllowWhitespace(msg)
fmt.Fprint(w, msg)
if !strings.HasSuffix(msg, "\n") {
w.Write([]byte("\n"))
}
}
func formatCertError(err error) string {
const unknownAuthority = `WARNING:
The proxy you are connecting to has presented a certificate signed by a
unknown authority. This is most likely due to either being presented
with a self-signed certificate or the certificate was truly signed by an
authority not known to the client.
If you know the certificate is self-signed and would like to ignore this
error use the --insecure flag.
If you have your own certificate authority that you would like to use to
validate the certificate chain presented by the proxy, set the
SSL_CERT_FILE and SSL_CERT_DIR environment variables respectively and try
again.
If you think something malicious may be occurring, contact your Teleport
system administrator to resolve this issue.
`
if errors.As(err, &x509.UnknownAuthorityError{}) {
return unknownAuthority
}
var hostnameErr x509.HostnameError
if errors.As(err, &hostnameErr) {
// Special case for connecting to Auth via Proxy using internal cluster domain.
if strings.HasSuffix(hostnameErr.Host, ".teleport.cluster.local") {
var proxyEnvBuilder strings.Builder
for _, key := range []string{
"https_proxy", "http_proxy", "no_proxy",
"HTTPS_PROXY", "HTTP_PROXY", "NO_PROXY",
} {
if val, ok := os.LookupEnv(key); ok {
fmt.Fprintf(&proxyEnvBuilder, " %s: %s\n", key, val)
}
}
return fmt.Sprintf(`Cannot connect to the Auth service via the Teleport Proxy.
There might be one or more network intermediaries (like a proxy or VPN) that are modifying your connection before it
reaches the Teleport Proxy. These intermediaries can alter how your connection is seen by the Teleport Proxy and
routed, leading to certificate mismatches.
To fix this, ensure that any network intermediaries are properly configured and not interfering with your connection.
DEBUG INFO:
Host: %s
Proxy Environment Variables:
%s
Server Certificate Details:
Subject: %s
Issuer: %s
Serial Number: %s
Not Before: %s
Not After: %s
DNS Names: %v
IP Addresses: %v`,
hostnameErr.Host,
proxyEnvBuilder.String(),
hostnameErr.Certificate.Subject,
hostnameErr.Certificate.Issuer,
hostnameErr.Certificate.SerialNumber,
hostnameErr.Certificate.NotBefore,
hostnameErr.Certificate.NotAfter,
hostnameErr.Certificate.DNSNames,
hostnameErr.Certificate.IPAddresses,
)
}
return fmt.Sprintf("Cannot establish https connection to %s:\n%s\n%s\n",
hostnameErr.Host,
hostnameErr.Error(),
"try a different hostname for --proxy or specify --insecure flag if you know what you're doing.")
}
var certInvalidErr x509.CertificateInvalidError
if errors.As(err, &certInvalidErr) {
return fmt.Sprintf(`WARNING:
The certificate presented by the proxy is invalid: %v.
Contact your Teleport system administrator to resolve this issue.`, certInvalidErr)
}
// Check for less explicit errors. These are often emitted on Darwin
if strings.Contains(err.Error(), "certificate is not trusted") {
return unknownAuthority
}
return ""
}
const (
// Bold is an escape code to format as bold or increased intensity
Bold = 1
// Red is an escape code for red terminal color
Red = 31
// Yellow is an escape code for yellow terminal color
Yellow = 33
// Blue is an escape code for blue terminal color
Blue = 36
// Gray is an escape code for gray terminal color
Gray = 37
)
// Color formats the string in a terminal escape color
func Color(color int, v any) string {
return fmt.Sprintf("\x1b[%dm%v\x1b[0m", color, v)
}
// InitCLIParser configures kingpin command line args parser with
// some defaults common for all Teleport CLI tools
func InitCLIParser(appName, appHelp string) (app *kingpin.Application) {
app = kingpin.New(appName, appHelp)
// make all flags repeatable, this makes the CLI easier to use.
app.AllRepeatable(true)
// hide "--help" flag
app.HelpFlag.Hidden()
app.HelpFlag.NoEnvar()
// write --help output to stdout instead of stderr
app.UsageWriter(os.Stdout)
return app.UsageRenderer(renderCompactUsage)
}
// InitHiddenCLIParser initializes a `kingpin.Application` that does not terminate the application
// or write any usage information to os.Stdout. Can be used in scenarios where multiple `kingpin.Application`
// instances are needed without interfering with subsequent parsing. Usage output is completely suppressed,
// and the default global `--help` flag is ignored to prevent the application from exiting.
func InitHiddenCLIParser() (app *kingpin.Application) {
app = kingpin.New("", "")
app.UsageWriter(io.Discard)
// HiddenHelpWriter suppresses flags like --completion-script-bash that override UsageWriter before writing.
// This prevents an unnecessary autocomplete from outputting.
app.HiddenHelpWriter(io.Discard)
app.Terminate(func(i int) {})
return app
}
// SplitIdentifiers splits list of identifiers by commas/spaces/newlines. Helpful when
// accepting lists of identifiers in CLI (role names, request IDs, etc).
func SplitIdentifiers(s string) []string {
return strings.FieldsFunc(s, func(r rune) bool {
return r == ',' || unicode.IsSpace(r)
})
}
// EscapeControl escapes all ANSI escape sequences from string and returns a
// string that is safe to print on the CLI. This is to ensure that malicious
// servers can not hide output. For more details, see:
// - https://sintonen.fi/advisories/scp-client-multiple-vulnerabilities.txt
func EscapeControl(s string) string {
if needsQuoting(s) {
return fmt.Sprintf("%q", s)
}
return s
}
// isAllowedWhitespace is a helper function for cli output escaping that returns
// true if a given rune is a whitespace character and allowed to be unescaped.
func isAllowedWhitespace(r rune) bool {
switch r {
case '\n', '\t', '\v':
// newlines, tabs, vertical tabs are allowed whitespace.
return true
}
return false
}
// AllowWhitespace escapes all ANSI escape sequences except some whitespace
// characters (\n \t \v) from string and returns a string that is safe to
// print on the CLI. This is to ensure that malicious servers can not hide
// output. For more details, see:
// - https://sintonen.fi/advisories/scp-client-multiple-vulnerabilities.txt
func AllowWhitespace(s string) string {
// loop over string searching for part to escape followed by allowed char.
// example: `\tabc\ndef\t\n`
// 1. part: "" sep: "\t"
// 2. part: "abc" sep: "\n"
// 3. part: "def" sep: "\t"
// 4. part: "" sep: "\n"
var sb strings.Builder
// note that increment also happens at bottom of loop because we can
// safely jump to place where allowedWhitespace was found.
for i := 0; i < len(s); i++ {
sepIdx := strings.IndexFunc(s[i:], isAllowedWhitespace)
if sepIdx == -1 {
// infalliable call, ignore error.
_, _ = sb.WriteString(EscapeControl(s[i:]))
// no separators remain.
break
}
part := EscapeControl(s[i : i+sepIdx])
_, _ = sb.WriteString(part)
sep := s[i+sepIdx]
_ = sb.WriteByte(sep)
i += sepIdx
}
return sb.String()
}
// needsQuoting returns true if any non-printable characters are found.
func needsQuoting(text string) bool {
for _, r := range text {
if !strconv.IsPrint(r) {
return true
}
}
return false
}
// defaultCommandPrintfWidth is the default command printf width.
const defaultCommandPrintfWidth = 12
// renderCompactUsage is a kingpin UsageRenderer used by all binaries to
// output usage text. It uses a kingpin UsageRenderer instead of
// UsageTemplate/UsageFuncs in order to avoid text/template (and reflect.MethodByName)
// so that the linker can enable dead-code elimination.
func renderCompactUsage(w io.Writer, ctx *kingpin.UsageContext) error {
width := ctx.Width
if ctx.Context.SelectedCommand != nil {
cmd := ctx.Context.SelectedCommand
fmt.Fprintf(w, "usage: %s %s", ctx.App.Name, cmd.String())
kingpin.WriteFormatUsage(w, cmd.FlagGroupModel, cmd.ArgGroupModel, cmd.CmdGroupModel, cmd.Help, width)
} else {
fmt.Fprintf(w, "Usage: %s", ctx.App.Name)
kingpin.WriteFormatUsage(w, ctx.App.FlagGroupModel, ctx.App.ArgGroupModel, ctx.App.CmdGroupModel, ctx.App.Help, width)
}
fmt.Fprintln(w)
if len(ctx.Context.Flags) > 0 {
fmt.Fprintln(w, "Flags:")
kingpin.FormatTwoColumns(w, ctx.Indent, 2, width, kingpin.FlagsToTwoColumns(ctx.Context.Flags))
fmt.Fprintln(w)
}
if len(ctx.Context.Args) > 0 {
fmt.Fprintln(w, "Args:")
kingpin.FormatTwoColumns(w, ctx.Indent, 2, width, kingpin.ArgsToTwoColumns(ctx.Context.Args))
fmt.Fprintln(w)
}
writeCommands := func(cmds []*kingpin.CmdModel) {
cmdWidth := defaultCommandPrintfWidth
for _, cmd := range cmds {
if cmd.Hidden {
continue
}
cmdWidth = max(cmdWidth, len(cmd.Name))
}
for _, cmd := range cmds {
if cmd.Hidden {
continue
}
fmt.Fprintf(w, " %-*s", cmdWidth, cmd.Name)
if cmd.Default {
fmt.Fprint(w, " (Default)")
}
fmt.Fprintf(w, " %s\n", cmd.Help)
}
}
if ctx.Context.SelectedCommand != nil {
if ctx.Context.SelectedCommand.CmdGroupModel != nil && len(ctx.Context.SelectedCommand.Commands) > 0 {
fmt.Fprintln(w, "Commands:")
writeCommands(ctx.Context.SelectedCommand.Commands)
fmt.Fprintln(w)
}
if len(ctx.Context.SelectedCommand.Aliases) > 0 {
fmt.Fprintln(w, "Aliases:")
for _, alias := range ctx.Context.SelectedCommand.Aliases {
fmt.Fprintln(w, alias)
}
fmt.Fprintln(w)
}
} else if ctx.App.CmdGroupModel != nil && len(ctx.App.Commands) > 0 {
fmt.Fprintln(w, "Commands:")
writeCommands(ctx.App.Commands)
fmt.Fprintln(w)
fmt.Fprintf(w, "Try '%s help [command]' to get help for a given command.\n\n", ctx.App.Name)
}
return nil
}
// IsPredicateError determines if the error is from failing to parse predicate expression
// by checking if the error as a string contains predicate keywords.
func IsPredicateError(err error) bool {
return strings.Contains(err.Error(), "predicate expression")
}
type PredicateError struct {
Err error
}
func (p PredicateError) Error() string {
return fmt.Sprintf("%s\nCheck syntax at https://goteleport.com/docs/reference/predicate-language/#resource-filtering", p.Err.Error())
}
// FormatAlert formats and colors the alert message if possible.
func FormatAlert(alert types.ClusterAlert) string {
// TODO(timothyb89): Due to complications with globally enabling +
// properly resetting Windows terminal ANSI processing, for now we just
// disable color output. Otherwise, raw ANSI escapes will be visible to
// users.
var buf bytes.Buffer
switch runtime.GOOS {
case constants.WindowsOS:
fmt.Fprint(&buf, alert.Spec.Message)
default:
switch alert.Spec.Severity {
case types.AlertSeverity_HIGH:
fmt.Fprint(&buf, Color(Red, alert.Spec.Message))
case types.AlertSeverity_MEDIUM:
fmt.Fprint(&buf, Color(Yellow, alert.Spec.Message))
default:
fmt.Fprint(&buf, alert.Spec.Message)
}
}
return buf.String()
}
// FilterArguments filters the input arguments, keeping only those defined in the provided `kingpin.ApplicationModel`.
// For example, if the model defines only one boolean flag `--insecure`, all other arguments in `args`
// will be excluded, and only the `--insecure` flag will remain.
func FilterArguments(args []string, model *kingpin.ApplicationModel) []string {
var result []string
for _, flag := range model.Flags {
for i := range args {
if strings.HasPrefix(args[i], fmt.Sprint("--", flag.Name, "=")) {
result = append(result, args[i])
break
}
if args[i] == fmt.Sprint("--", flag.Name) {
if flag.IsBoolFlag() {
result = append(result, args[i])
} else if i+2 <= len(args) {
result = append(result, args[i], args[i+1])
}
break
}
}
}
return result
}
/*
* 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 utils
import (
"context"
"log/slog"
"os"
"path/filepath"
"strings"
"github.com/gravitational/trace"
)
// TryReadValueAsFile is a utility function to read a value
// from the disk if it looks like an absolute path,
// otherwise, treat it as a value.
// It only support absolute paths to avoid ambiguity in interpretation of the value
func TryReadValueAsFile(value string) (string, error) {
if !filepath.IsAbs(value) {
return value, nil
}
// treat it as an absolute filepath
contents, err := os.ReadFile(value)
if err != nil {
return "", trace.ConvertSystemError(err)
}
// trim newlines as tokens in files tend to have newlines
out := strings.TrimSpace(string(contents))
if out == "" {
slog.WarnContext(context.Background(), "Empty config value file", "file", value)
}
return out, nil
}
/*
* 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 utils
import (
"io"
"net"
"sync"
"sync/atomic"
"github.com/gravitational/trace"
)
// NewWaitConn returns new connection wrapper that
// provides the ability to wait for the connection
// to be closed.
func NewWaitConn(conn net.Conn) *WaitConn {
wc := &WaitConn{
Conn: conn,
closed: make(chan struct{}),
}
wc.close = sync.OnceValue(func() error {
err := wc.Conn.Close()
close(wc.closed)
return trace.Wrap(err)
})
return wc
}
// WaitConn wraps a connection and provides the ability to wait for the
// connection to be closed.
type WaitConn struct {
net.Conn
close func() error
closed chan struct{}
}
// Close closes the connection.
func (c *WaitConn) Close() error {
return c.close()
}
func (c *WaitConn) Done() <-chan struct{} { return c.closed }
func (c *WaitConn) Wait() { <-c.Done() }
// TrackingConn is a net.Conn that keeps track of how much data was transmitted
// (TX) and received (RX) over the net.Conn. A maximum of about 18446
// petabytes can be kept track of for TX and RX before it rolls over.
// See https://golang.org/ref/spec#Numeric_types for more details.
type TrackingConn struct {
net.Conn
r *trackingReader
w *trackingWriter
}
// NewTrackingConn returns a net.Conn that can keep track of how much data was
// transmitted over it.
func NewTrackingConn(conn net.Conn) *TrackingConn {
return &TrackingConn{
Conn: conn,
r: &trackingReader{r: conn},
w: &trackingWriter{w: conn},
}
}
// Stat returns the transmitted (TX) and received (RX) bytes over the net.Conn.
func (s *TrackingConn) Stat() (written, read uint64) {
return s.w.Count(), s.r.Count()
}
func (s *TrackingConn) Read(b []byte) (n int, err error) {
return s.r.Read(b)
}
func (s *TrackingConn) Write(b []byte) (n int, err error) {
return s.w.Write(b)
}
// trackingReader is an io.Reader that counts the total number of bytes read.
// It's thread-safe if the underlying io.Reader is thread-safe.
type trackingReader struct {
r io.Reader
count uint64
}
// Count returns the total number of bytes read so far.
func (r *trackingReader) Count() uint64 {
return atomic.LoadUint64(&r.count)
}
func (r *trackingReader) Read(b []byte) (int, error) {
n, err := r.r.Read(b)
atomic.AddUint64(&r.count, uint64(n))
// This has to use the original error type or else utilities using the connection
// (like io.Copy, which is used by the oxy forwarder) may incorrectly categorize
// the error produced by this and terminate the connection unnecessarily.
return n, err
}
// trackingWriter is an io.Writer that counts the total number of bytes
// written.
// It's thread-safe if the underlying io.Writer is thread-safe.
type trackingWriter struct {
count uint64 // intentionally placed first to ensure 64-bit alignment
w io.Writer
}
// Count returns the total number of bytes written so far.
func (w *trackingWriter) Count() uint64 {
return atomic.LoadUint64(&w.count)
}
func (w *trackingWriter) Write(b []byte) (int, error) {
n, err := w.w.Write(b)
atomic.AddUint64(&w.count, uint64(n))
return n, trace.Wrap(err)
}
// ConnWithAddr is a [net.Conn] wrapper that allows the local and remote address
// to be overridden.
type ConnWithAddr struct {
net.Conn
localAddrOverride net.Addr
remoteAddrOverride net.Addr
}
// LocalAddr implements [net.Conn].
func (c *ConnWithAddr) LocalAddr() net.Addr {
if c.localAddrOverride != nil {
return c.localAddrOverride
}
return c.Conn.LocalAddr()
}
// RemoteAddr implements [net.Conn].
func (c *ConnWithAddr) RemoteAddr() net.Addr {
if c.remoteAddrOverride != nil {
return c.remoteAddrOverride
}
return c.Conn.RemoteAddr()
}
// NetConn returns the underlying [net.Conn].
func (c *ConnWithAddr) NetConn() net.Conn {
return c.Conn
}
// NewConnWithSrcAddr wraps provided connection and overrides client remote address.
func NewConnWithSrcAddr(conn net.Conn, clientSrcAddr net.Addr) *ConnWithAddr {
return &ConnWithAddr{
Conn: conn,
remoteAddrOverride: clientSrcAddr,
}
}
// NewConnWithAddr wraps a [net.Conn] optionally overriding the local and remote
// addresses with the provided ones, if non-nil.
func NewConnWithAddr(conn net.Conn, localAddr, remoteAddr net.Addr) *ConnWithAddr {
return &ConnWithAddr{
Conn: conn,
localAddrOverride: localAddr,
remoteAddrOverride: remoteAddr,
}
}
/*
* 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 utils
// ReplaceInSlice replaces element old with new and returns a new slice.
func ReplaceInSlice(s []string, old string, new string) []string {
out := make([]string, 0, len(s))
for _, x := range s {
if x == old {
out = append(out, new)
} else {
out = append(out, x)
}
}
return out
}
//go:build !windows
/*
* 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 utils
import (
"errors"
"io/fs"
"os"
"os/user"
"path/filepath"
"strconv"
"syscall"
"github.com/gravitational/trace"
hostuser "github.com/gravitational/teleport/session/host/user"
)
// PercentUsed returns percentage of disk space used. The percentage of disk
// space used is calculated from (total blocks - free blocks)/total blocks.
// The value is rounded to the nearest whole integer.
func PercentUsed(path string) (float64, error) {
var stat syscall.Statfs_t
err := syscall.Statfs(path, &stat)
if err != nil {
return 0, trace.Wrap(err)
}
ratio := float64(stat.Blocks-stat.Bfree) / float64(stat.Blocks)
return Round(ratio * 100), nil
}
// FreeDiskWithReserve returns the available disk space (in bytes) on the disk at dir, minus `reservedFreeDisk`.
func FreeDiskWithReserve(dir string, reservedFreeDisk uint64) (uint64, error) {
var stat syscall.Statfs_t
err := syscall.Statfs(dir, &stat)
if err != nil {
return 0, trace.Wrap(err)
}
//nolint:unconvert // The cast is only necessary for linux platform.
avail := uint64(stat.Bavail) * uint64(stat.Bsize)
if reservedFreeDisk > avail {
return 0, trace.Errorf("no free space left")
}
return avail - reservedFreeDisk, nil
}
// CanUserWriteTo attempts to check if a user has write access to certain path.
// It also works around the program being run as root and tries to check
// the permissions of the user who executed the program as root.
// This should only be used for string formatting or inconsequential use cases
// as it's not bullet proof and can report wrong results.
func CanUserWriteTo(path string) (bool, error) {
// prevent infinite loops with a max dir depth
var fileInfo os.FileInfo
var err error
for range 20 {
fileInfo, err = os.Stat(path)
if err == nil {
break
}
if errors.Is(err, fs.ErrNotExist) {
path = filepath.Dir(path)
continue
}
return false, trace.BadParameter("failed to find path: %+v", err)
}
var uid int
var gid int
if stat, ok := fileInfo.Sys().(*syscall.Stat_t); ok {
uid = int(stat.Uid)
gid = int(stat.Gid)
}
var usr *user.User
if ogUser := os.Getenv("SUDO_USER"); ogUser != "" {
usr, err = hostuser.Lookup(ogUser)
if err != nil {
return false, trace.NotFound("could not determine original user: %+v", err)
}
} else {
usr, err = hostuser.Current()
if err != nil {
return false, trace.NotFound("could not determine current user: %+v", err)
}
}
perm := fileInfo.Mode().Perm()
// file is owned by the user
if strconv.Itoa(uid) == usr.Uid {
// file has u+wx permissions
if perm&syscall.S_IWUSR != 0 &&
perm&syscall.S_IXUSR != 0 {
return true, nil
}
}
// file and user have a group in common
groupIDs, err := hostuser.GroupIds(usr)
if err != nil {
return false, trace.NotFound("could not determine current user group ids: %+v", err)
}
for _, groupID := range groupIDs {
if strconv.Itoa(gid) == groupID {
// file has g+wx permissions
if perm&syscall.S_IWGRP != 0 &&
perm&syscall.S_IXGRP != 0 {
return true, nil
}
break
}
}
// file has o+wx permissions
if perm&syscall.S_IWOTH != 0 &&
perm&syscall.S_IXOTH != 0 {
return true, nil
}
return false, nil
}
/*
* 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 utils
// StringMapsEqual returns true if two strings maps are equal
func StringMapsEqual(a, b map[string]string) bool {
if len(a) != len(b) {
return false
}
for key := range a {
if a[key] != b[key] {
return false
}
}
return true
}
/*
* 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 utils
import (
"context"
"errors"
"io"
"net"
"strings"
"syscall"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/constants"
)
// IsUseOfClosedNetworkError returns true if the specified error
// indicates the use of a closed network connection.
func IsUseOfClosedNetworkError(err error) bool {
if err == nil {
return false
}
return errors.Is(err, net.ErrClosed) || strings.Contains(err.Error(), constants.UseOfClosedNetworkConnection)
}
// IsFailedToSendCloseNotifyError returns true if the provided error is the
// "tls: failed to send closeNotify".
func IsFailedToSendCloseNotifyError(err error) bool {
if err == nil {
return false
}
return strings.Contains(err.Error(), constants.FailedToSendCloseNotify)
}
// IsOKNetworkError returns true if the provided error received from a network
// operation is one of those that usually indicate normal connection close. If
// the error is a trace.Aggregate, all the errors must be OK network errors.
func IsOKNetworkError(err error) bool {
// trace.Aggregate contains at least one error and all the errors are
// non-nil
var a trace.Aggregate
if errors.As(trace.Unwrap(err), &a) {
for _, err := range a.Errors() {
if !IsOKNetworkError(err) {
return false
}
}
return true
}
return errors.Is(err, io.EOF) || IsUseOfClosedNetworkError(err) || IsFailedToSendCloseNotifyError(err)
}
// IsConnectionRefused returns true if the given err is "connection refused" error.
func IsConnectionRefused(err error) bool {
var errno syscall.Errno
if errors.As(err, &errno) {
return errors.Is(errno, syscall.ECONNREFUSED)
}
return false
}
// IsConnectionError returns true when the error indicates a socket-level connection failure
// (e.g. file not found, connection refused), as opposed to an HTTP or application-level error.
func IsConnectionError(err error) bool {
var opErr *net.OpError
return errors.As(err, &opErr)
}
// IsUntrustedCertErr checks if an error is an untrusted cert error.
func IsUntrustedCertErr(err error) bool {
if err == nil {
return false
}
errMsg := err.Error()
return strings.Contains(errMsg, "x509: certificate is valid for") ||
strings.Contains(errMsg, "certificate is not trusted")
}
// CanExplainNetworkError returns a simple to understand error message that can
// be used to debug common network and/or protocol errors.
func CanExplainNetworkError(err error) (string, bool) {
var (
derr *net.DNSError
nerr net.Error
)
switch {
// Connection refused errors can be reproduced by attempting to connect to a
// host:port that no process is listening on. The raw error typically looks
// like the following:
//
// dial tcp 127.0.0.1:8000: connect: connection refused
case errors.Is(err, syscall.ECONNREFUSED):
return `Connection Refused
Teleport was unable to connect to the requested host, possibly because the server is not running. Ensure the server is running and listening on the correct port.
Use "nc -vz HOST PORT" to help debug this issue.`, true
// Host unreachable errors can be reproduced by running
// "ip route add unreachable HOST" to update the routing table to make
// the host unreachable. Packets will be discarded and an ICMP message
// will be returned. The raw error typically looks like the following:
//
// dial tcp 10.10.10.10:8000: connect: no route to host
case errors.Is(err, syscall.EHOSTUNREACH):
return `No Route to Host
Teleport could not connect to the requested host, likely because there is no valid network path to reach it. Check the network routing table to ensure a valid path to the host exists.
Use "ping HOST" and "ip route get HOST" to help debug this issue.`, true
// Connection reset errors can be reproduced by creating a HTTP server that
// accepts requests but closes the connection before writing a response. The
// raw error typically looks like the following:
//
// read tcp 127.0.0.1:49764->127.0.0.1:8000: read: connection reset by peer
case errors.Is(err, syscall.ECONNRESET):
return `Connection Reset by Peer
Teleport could not complete the request because the server abruptly closed the connection before the response was received. To resolve this issue, ensure the server (or load balancer) does not have a timeout terminating the connection early and verify that the server is not crash looping.
Use protocol-specific tools (e.g., curl, psql) to help debug this issue.`, true
// Slow responses can be reproduced by creating a HTTP server that
// does a time.Sleep before responding. The raw error typically
// looks like the following:
//
// context deadline exceeded
//
// HTTP/1 wraps context.DeadlineExceeded via timeoutError.Is.
// HTTP/2's http2httpError does not wrap it; it implements
// net.Error with Timeout() returning true. Match both shapes
// here, plus syscall.ETIMEDOUT and other net.Error timeouts.
case errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &nerr) && nerr.Timeout()):
return `Context Deadline Exceeded
Teleport did not receive a response within the timeout period, likely due to the system being overloaded, network congestion, or a firewall blocking traffic. To resolve this issue, connect to the host directly and ensure it is responding promptly.
Use protocol-specific tools (e.g., curl, psql) to assist in debugging this issue.`, true
// No such host errors can be reproduced by attempting to resolve a invalid
// domain name. The raw error typically looks like the following:
//
// dial tcp: lookup qweqweqwe.com: no such host
case errors.As(err, &derr) && derr.IsNotFound:
return `No Such Host
Teleport was unable to resolve the provided domain name, likely because the domain does not exist. To resolve this issue, verify the domain is correct and ensure the DNS resolver is properly resolving it.
Use "dig +short HOST" to help debug this issue.`, true
}
return "", false
}
const (
// SelfSignedCertsMsg is a helper message to point users towards helpful documentation.
SelfSignedCertsMsg = "Your proxy certificate is not trusted or expired. " +
"Please update the certificate or follow this guide for self-signed certs: https://goteleport.com/docs/admin-guides/management/admin/self-signed-certs/"
)
/*
* 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 utils
import (
"slices"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
setutils "github.com/gravitational/teleport/lib/utils/set"
)
// Fields represents a generic string-keyed map.
type Fields map[string]any
// GetString returns a string representation of a field.
func (f Fields) GetString(key string) string {
val, found := f[key]
if !found {
return ""
}
return val.(string)
}
// GetStrings returns a slice-of-strings representation of a field.
func (f Fields) GetStrings(key string) []string {
val, found := f[key]
if !found {
return nil
}
res, _ := getStrings(val)
return res
}
func getStrings(val any) ([]string, bool) {
strings, ok := val.([]string)
if ok {
return strings, true
}
slice, ok := val.([]any)
if !ok {
return nil, false
}
res := make([]string, 0, len(slice))
for _, v := range slice {
s, ok := v.(string)
if ok {
res = append(res, s)
}
}
return res, true
}
// GetInt returns an int representation of a field.
func (f Fields) GetInt(key string) int {
val, found := f[key]
if !found {
return 0
}
v, ok := val.(int)
if !ok {
f, ok := val.(float64)
if ok {
v = int(f)
}
}
return v
}
// GetTime returns a time.Time representation of a field.
func (f Fields) GetTime(key string) time.Time {
val, found := f[key]
if !found {
return time.Time{}
}
v, ok := val.(time.Time)
if !ok {
s := f.GetString(key)
v, _ = time.Parse(time.RFC3339, s)
}
return v
}
// HasField returns true if the field exists.
func (f Fields) HasField(key string) bool {
_, ok := f[key]
return ok
}
func (f Fields) Get(key string) (any, bool) {
v, ok := f[key]
return v, ok
}
func (f Fields) GetMapEntry(mapRef *types.WhereExpr2) (any, bool) {
field := mapRef.L.Field
key, ok := mapRef.R.Literal.(string)
if !ok {
return nil, false
}
val, found := f.Get(field)
if !found {
return nil, false
}
m, ok := val.(map[string]any)
if !ok {
return nil, false
}
v, ok := m[key]
if !ok {
return nil, false
}
return v, true
}
// FieldsCondition is a boolean function on Fields.
type FieldsCondition func(Fields) bool
// ToFieldsConditionConfig is the configuration for ToFieldsCondition.
type ToFieldsConditionConfig struct {
Expr *types.WhereExpr
// CanView is an optional function that checks if the user is allowed to view the resource.
CanView func(Fields) bool
}
// ToFieldsCondition converts a WhereExpr into a FieldsCondition.
func ToFieldsCondition(cfg ToFieldsConditionConfig) (FieldsCondition, error) {
expr := cfg.Expr
if cfg.Expr == nil {
return nil, trace.BadParameter("expr is nil")
}
binOp := func(e types.WhereExpr2, op func(a, b bool) bool) (FieldsCondition, error) {
left, err := ToFieldsCondition(ToFieldsConditionConfig{
Expr: e.L,
CanView: cfg.CanView,
})
if err != nil {
return nil, trace.Wrap(err)
}
right, err := ToFieldsCondition(
ToFieldsConditionConfig{
Expr: e.R,
CanView: cfg.CanView,
})
if err != nil {
return nil, trace.Wrap(err)
}
return func(f Fields) bool { return op(left(f), right(f)) }, nil
}
if expr, err := binOp(expr.And, func(a, b bool) bool { return a && b }); err == nil {
return expr, nil
}
if expr, err := binOp(expr.Or, func(a, b bool) bool { return a || b }); err == nil {
return expr, nil
}
if inner, err := ToFieldsCondition(
ToFieldsConditionConfig{
Expr: expr.Not,
CanView: cfg.CanView,
},
); err == nil {
return func(f Fields) bool { return !inner(f) }, nil
}
if expr.Equals.L != nil && expr.Equals.R != nil {
left, right := expr.Equals.L, expr.Equals.R
switch {
case left.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(left.MapRef)
if !ok {
return false
}
strs, ok := val.(string)
if !ok {
return false
}
var strsVals string
if right.Field != "" {
strsVals = f.GetString(right.Field)
} else if right.Literal != nil {
strsVals, ok = right.Literal.(string)
if !ok {
return false
}
}
return strs == strsVals
}, nil
case right.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(right.MapRef)
if !ok {
return false
}
str, ok := val.(string)
if !ok {
return false
}
var strsVals string
if left.Field != "" {
strsVals = f.GetString(left.Field)
} else if left.Literal != nil {
strsVals, ok = left.Literal.(string)
if !ok {
return false
}
}
return str == strsVals
}, nil
case left.Field != "" && right.Field != "":
return func(f Fields) bool { return f[left.Field] == f[right.Field] }, nil
case left.Literal != nil && right.Field != "":
return func(f Fields) bool { return left.Literal == f[right.Field] }, nil
case left.Field != "" && right.Literal != nil:
return func(f Fields) bool { return f[left.Field] == right.Literal }, nil
}
}
if expr.Contains.L != nil && expr.Contains.R != nil {
left, right := expr.Contains.L, expr.Contains.R
switch {
case left.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(left.MapRef)
if !ok {
return false
}
strs, ok := getStrings(val)
if !ok {
return false
}
var strsVals string
if right.Field != "" {
strsVals = f.GetString(right.Field)
} else if right.Literal != nil {
strsVals, ok = right.Literal.(string)
if !ok {
return false
}
}
return slices.Contains(strs, strsVals)
}, nil
case right.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(right.MapRef)
if !ok {
return false
}
str, ok := val.(string)
if !ok {
return false
}
var strsVals []string
if left.Field != "" {
strsVals = f.GetStrings(left.Field)
} else if left.Literal != nil {
strsVals, ok = getStrings(left.Literal)
if !ok {
return false
}
}
return slices.Contains(strsVals, str)
}, nil
case left.Field != "" && right.Field != "":
return func(f Fields) bool { return slices.Contains(f.GetStrings(left.Field), f.GetString(right.Field)) }, nil
case left.Literal != nil && right.Field != "":
if ss, ok := getStrings(left.Literal); ok {
return func(f Fields) bool { return slices.Contains(ss, f.GetString(right.Field)) }, nil
}
case left.Field != "" && right.Literal != nil:
if s, ok := right.Literal.(string); ok {
return func(f Fields) bool { return slices.Contains(f.GetStrings(left.Field), s) }, nil
}
}
}
if expr.ContainsAll.L != nil && expr.ContainsAll.R != nil {
left, right := expr.ContainsAll.L, expr.ContainsAll.R
switch {
case left.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(left.MapRef)
if !ok {
return false
}
strs, ok := getStrings(val)
if !ok {
return false
}
var strsVals []string
if right.Field != "" {
strsVals = f.GetStrings(right.Field)
} else if right.Literal != nil {
strsVals, ok = getStrings(right.Literal)
if !ok {
return false
}
}
return containsAll(strs, strsVals)
}, nil
case right.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(right.MapRef)
if !ok {
return false
}
strs, ok := getStrings(val)
if !ok {
return false
}
var strsVals []string
if left.Field != "" {
strsVals = f.GetStrings(left.Field)
} else if left.Literal != nil {
strsVals, ok = getStrings(left.Literal)
if !ok {
return false
}
}
return containsAll(strsVals, strs)
}, nil
case left.Field != "" && right.Field != "":
return func(f Fields) bool { return containsAll(f.GetStrings(left.Field), f.GetStrings(right.Field)) }, nil
case left.Literal != nil && right.Field != "":
if ss, ok := getStrings(left.Literal); ok {
return func(f Fields) bool { return containsAll(ss, f.GetStrings(right.Field)) }, nil
}
case left.Field != "" && right.Literal != nil:
if s, ok := getStrings(right.Literal); ok {
return func(f Fields) bool { return containsAll(f.GetStrings(left.Field), s) }, nil
}
}
}
if expr.ContainsAny.L != nil && expr.ContainsAny.R != nil {
left, right := expr.ContainsAny.L, expr.ContainsAny.R
switch {
case left.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(left.MapRef)
if !ok {
return false
}
strs, ok := getStrings(val)
if !ok {
return false
}
var strsVals []string
if right.Field != "" {
strsVals = f.GetStrings(right.Field)
} else if right.Literal != nil {
strsVals, ok = getStrings(right.Literal)
if !ok {
return false
}
}
return containsAny(strs, strsVals)
}, nil
case right.MapRef != nil:
return func(f Fields) bool {
val, ok := f.GetMapEntry(right.MapRef)
if !ok {
return false
}
strs, ok := getStrings(val)
if !ok {
return false
}
var strsVals []string
if left.Field != "" {
strsVals = f.GetStrings(left.Field)
} else if left.Literal != nil {
strsVals, ok = getStrings(left.Literal)
if !ok {
return false
}
}
return containsAny(strsVals, strs)
}, nil
case left.Field != "" && right.Field != "":
return func(f Fields) bool { return containsAny(f.GetStrings(left.Field), f.GetStrings(right.Field)) }, nil
case left.Literal != nil && right.Field != "":
if ss, ok := getStrings(left.Literal); ok {
return func(f Fields) bool { return containsAny(ss, f.GetStrings(right.Field)) }, nil
}
case left.Field != "" && right.Literal != nil:
if s, ok := getStrings(right.Literal); ok {
return func(f Fields) bool { return containsAny(f.GetStrings(left.Field), s) }, nil
}
}
}
if expr.CanView != nil {
if cfg.CanView == nil {
return nil, trace.BadParameter("canView expression provided but no canView function specified")
}
return func(f Fields) bool { return cfg.CanView(f) }, nil
}
return nil, trace.BadParameter("failed to convert expression %q to FieldsCondition", expr)
}
func containsAll(slice []string, items []string) bool {
set := setutils.New(slice...)
if len(items) == 0 {
return false
}
for _, item := range items {
if !set.Contains(item) {
return false
}
}
return true
}
func containsAny(slice []string, items []string) bool {
set := setutils.New(slice...)
if len(items) == 0 {
return false
}
for _, item := range items {
if set.Contains(item) {
return true
}
}
return false
}
/*
* 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 utils
import (
"context"
"errors"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
oteltrace "go.opentelemetry.io/otel/trace"
)
// NOTE: when making changes to this file, run tests with `TEST_FNCACHE_FUZZY=yes` to enable
// additional fuzzy tests which aren't run during normal CI.
// ErrFnCacheClosed is returned from Get when the FnCache context is closed
var ErrFnCacheClosed = errors.New("fncache permanently closed")
// FnCache is a helper for temporarily storing the results of regularly called functions. This helper is
// used to limit the amount of backend reads that occur while the primary cache is unhealthy. Most resources
// do not require this treatment, however, certain resources (cas, nodes, etc.) can be loaded on a per-request
// basis and can cause a significant number of backend reads if the cache is unhealthy or taking a while to initialize.
type FnCache struct {
cfg FnCacheConfig
closed bool
cancel context.CancelFunc
nextCleanup time.Time
mu sync.Mutex
entries map[any]*fnCacheEntry
}
// cleanupMultiplier is an arbitrary multiplier used to derive the default interval
// for periodic lazy cleanup of expired entries. This cache is typically used to
// store a small number of regularly read keys, so most old values aught to be
// removed upon subsequent reads of the same key. If the cache is being used in a
// context where keys might become regularly orphaned (no longer read), then a
// custom CleanupInterval should be provided.
const cleanupMultiplier time.Duration = 16
// FnCacheConfig contains dependencies for a FnCache.
type FnCacheConfig struct {
// TTL is the time to live for cache entries.
TTL time.Duration
// Clock is the clock used to determine the current time.
Clock clockwork.Clock
// Context is the context used to cancel the cache. All loadfns
// will be provided with this context.
Context context.Context
// ReloadOnErr causes entries to be reloaded immediately if
// the currently loaded value is an error. Note that all concurrent
// requests registered before load completes still observe the
// same error. This option is only really useful for longer TTLs.
ReloadOnErr bool
// CleanupInterval is the interval at which cleanups occur (defaults to
// 16x the supplied TTL). Longer cleanup intervals are appropriate for
// caches where keys are unlikely to become orphaned. Shorter cleanup
// intervals should be used when keys regularly become orphaned.
CleanupInterval time.Duration
// OnExpiry is an optional callback that will be executed any time an
// item is expired and removed from the cache, or replaced by a reload
// in get() when the entry's TTL has elapsed. The callback must not call
// any method on the same FnCache instance because the cache mutex may
// be held when the callback is invoked.
OnExpiry func(ctx context.Context, key any, value any)
}
// CheckAndSetDefaults validates the FnCacheConfig is populated
// with required fields and sets any omitted fields to default values.
func (c *FnCacheConfig) CheckAndSetDefaults() error {
if c.TTL <= 0 {
return trace.BadParameter("missing TTL parameter")
}
if c.Clock == nil {
c.Clock = clockwork.NewRealClock()
}
if c.Context == nil {
c.Context = context.Background()
}
if c.CleanupInterval <= 0 {
c.CleanupInterval = c.TTL * cleanupMultiplier
}
return nil
}
// NewFnCache creates a [FnCache] from the provided [FnCacheConfig].
func NewFnCache(cfg FnCacheConfig) (*FnCache, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
cache := &FnCache{
cfg: cfg,
entries: make(map[any]*fnCacheEntry),
}
cache.cfg.Context, cache.cancel = context.WithCancel(cfg.Context)
return cache, nil
}
type fnCacheEntry struct {
v any
e error
t time.Time
ttl time.Duration
loaded chan struct{}
}
// Shutdown expires all items in the cache. If the OnExpiry
// callback was set in the FnCacheConfig it will be called once
// per item in the cache.
func (c *FnCache) Shutdown(ctx context.Context) {
c.mu.Lock()
c.cancel()
c.closed = true
entries := c.entries
c.entries = make(map[any]*fnCacheEntry)
c.mu.Unlock()
// non-blocking eviction
for key, entry := range entries {
select {
case <-entry.loaded:
if c.cfg.OnExpiry != nil && entry.e == nil {
c.cfg.OnExpiry(ctx, key, entry.v)
}
delete(entries, key)
case <-ctx.Done():
return
default:
// entry is still being loaded
}
}
// blocking eviction
for key, entry := range entries {
select {
case <-entry.loaded:
if c.cfg.OnExpiry != nil && entry.e == nil {
c.cfg.OnExpiry(ctx, key, entry.v)
}
delete(entries, key)
case <-ctx.Done():
return
}
}
}
// Remove purges a specific item in the cache.
func (c *FnCache) Remove(key any) {
c.mu.Lock()
defer c.mu.Unlock()
delete(c.entries, key)
}
// Set places an item in the cache using the default TTL.
func (c *FnCache) Set(key, value any) {
c.SetWithTTL(key, value, c.cfg.TTL)
}
// GetIfExists retrieves a value from the cache without triggering a load operation.
// It returns (value, true) if a valid, non-expired entry exists, or (nil, false)
// otherwise. If an entry is currently being loaded by FnCacheGet, Get will
// return false immediately without blocking. Get returns false for entries that
// contain errors.
// For most of the cases the FnCacheGet function should be used instead.
func (c *FnCache) GetIfExists(key any) (any, bool) {
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return nil, false
}
entry := c.entries[key]
c.mu.Unlock()
if entry == nil {
return nil, false
}
select {
case <-entry.loaded:
if c.cfg.Clock.Now().After(entry.t.Add(entry.ttl)) {
return nil, false
}
if entry.e != nil {
return nil, false
}
return entry.v, true
default:
// Entry still loading - treat as cache miss
return nil, false
}
}
// SetWithTTL places an item in the cache with an explicit TTL.
func (c *FnCache) SetWithTTL(key, value any, ttl time.Duration) {
c.mu.Lock()
defer c.mu.Unlock()
loaded := make(chan struct{})
close(loaded)
c.entries[key] = &fnCacheEntry{
v: value,
t: c.cfg.Clock.Now(),
ttl: ttl,
loaded: loaded,
}
}
// RemoveExpired purges any items from the cache which have exceeded their TTL.
func (c *FnCache) RemoveExpired() {
c.mu.Lock()
defer c.mu.Unlock()
now := c.cfg.Clock.Now()
c.removeExpiredLocked(now)
c.nextCleanup = now.Add(c.cfg.CleanupInterval)
}
func (c *FnCache) removeExpiredLocked(now time.Time) {
for key, entry := range c.entries {
select {
case <-entry.loaded:
if now.After(entry.t.Add(entry.ttl)) {
if c.cfg.OnExpiry != nil && entry.e == nil {
c.cfg.OnExpiry(context.WithoutCancel(c.cfg.Context), key, entry.v)
}
delete(c.entries, key)
}
default:
// entry is still being loaded
}
}
}
// FnCacheGet loads the result associated with the supplied key. If no result is currently stored, or the stored result
// was acquired >TTL ago, then loadfn is used to reload it. Subsequent calls while the value is being loaded/reloaded
// block until the first call updates the entry. Note that the supplied context can cancel the call to Get, but will
// not cancel loading. The supplied loadfn should not be canceled just because the specific request happens to have
// been canceled.
func FnCacheGet[K comparable, T any](ctx context.Context, cache *FnCache, key K, loadfn func(ctx context.Context) (T, error)) (T, error) {
return FnCacheGetWithTTL(ctx, cache, key, cache.cfg.TTL, loadfn)
}
// FnCacheGetWithTTL is identical to FnCacheGet except that it allows individual keys to specify
// a TTL that is used instead of the configured TTL for the FnCache.
func FnCacheGetWithTTL[K comparable, T any](ctx context.Context, cache *FnCache, key K, ttl time.Duration, loadfn func(ctx context.Context) (T, error)) (T, error) {
t, err := cache.get(ctx, key, ttl, func(ctx context.Context) (any, error) {
return loadfn(ctx)
})
ret, ok := t.(T)
switch {
case err != nil:
return ret, err
case t == nil:
return ret, nil
case !ok:
return ret, trace.BadParameter("value retrieved was %T, expected %T", t, ret)
}
return ret, err
}
// get loads the result associated with the supplied key. If no result is currently stored, or the stored result
// was acquired >ttl ago, then loadfn is used to reload it. Subsequent calls while the value is being loaded/reloaded
// block until the first call updates the entry. Note that the supplied context can cancel the call to Get, but will
// not cancel loading. The supplied loadfn should not be canceled just because the specific request happens to have
// been canceled.
func (c *FnCache) get(ctx context.Context, key any, ttl time.Duration, loadfn func(ctx context.Context) (any, error)) (any, error) {
select {
case <-c.cfg.Context.Done():
return nil, ErrFnCacheClosed
default:
}
c.mu.Lock()
if c.closed {
c.mu.Unlock()
return nil, ErrFnCacheClosed
}
now := c.cfg.Clock.Now()
// Check if we need to perform periodic cleanup.
if now.After(c.nextCleanup) {
c.removeExpiredLocked(now)
c.nextCleanup = now.Add(c.cfg.CleanupInterval)
}
entry := c.entries[key]
needsReload := true
if entry != nil {
select {
case <-entry.loaded:
needsReload = now.After(entry.t.Add(entry.ttl))
if c.cfg.ReloadOnErr && entry.e != nil {
needsReload = true
}
default:
// reload is already in progress
needsReload = false
}
}
if needsReload {
// If we are replacing a loaded, successful entry, call OnExpiry so
// the old value is properly cleaned up. Without this, entries that
// expire between cleanup intervals are silently dropped when a new
// request triggers a reload, skipping the OnExpiry callback.
if entry != nil && entry.e == nil && c.cfg.OnExpiry != nil {
c.cfg.OnExpiry(context.WithoutCancel(c.cfg.Context), key, entry.v)
}
// Insert a new entry with a new loaded channel. This channel will
// block subsequent reads, and serve as a memory barrier for the results.
entry = &fnCacheEntry{
loaded: make(chan struct{}),
ttl: ttl,
}
c.entries[key] = entry
go func() {
// Link the config context with the span from ctx, if one exists,
// so that the loadfn can be traced appropriately.
loadCtx := oteltrace.ContextWithSpan(c.cfg.Context, oteltrace.SpanFromContext(ctx))
entry.v, entry.e = loadfn(loadCtx)
entry.t = c.cfg.Clock.Now()
close(entry.loaded)
}()
}
c.mu.Unlock()
// Wait for the result to be loaded (this is also a memory barrier).
select {
case <-entry.loaded:
return entry.v, entry.e
case <-ctx.Done():
return nil, ctx.Err()
case <-c.cfg.Context.Done():
return nil, ErrFnCacheClosed
}
}
/*
* 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 utils
import (
"context"
"crypto/rand"
"errors"
"io"
"io/fs"
"os"
"path/filepath"
"runtime"
"strings"
"syscall"
"time"
"github.com/gofrs/flock"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
)
// ErrUnsuccessfulLockTry designates an error when we temporarily couldn't acquire lock
// (most probably it was already locked by someone else), another try might succeed.
var ErrUnsuccessfulLockTry = errors.New("could not acquire lock on the file at this time")
const (
// FSLockRetryDelay is a delay between attempts to acquire lock.
FSLockRetryDelay = 10 * time.Millisecond
)
// OpenFileWithFlagsFunc defines a function used to open files providing options.
type OpenFileWithFlagsFunc func(name string, flag int, perm os.FileMode) (*os.File, error)
// EnsureLocalPath makes sure the path exists, or, if omitted results in the subpath in
// default gravity config directory, e.g.
//
// EnsureLocalPath("/custom/myconfig", ".gravity", "config") -> /custom/myconfig
// EnsureLocalPath("", ".gravity", "config") -> ${HOME}/.gravity/config
//
// It also makes sure that base dir exists
func EnsureLocalPath(customPath string, defaultLocalDir, defaultLocalPath string) (string, error) {
if customPath == "" {
homeDir, err := os.UserHomeDir()
if err != nil || homeDir == "" {
return "", trace.BadParameter("could not get user home dir: %v", err)
}
customPath = filepath.Join(homeDir, defaultLocalDir, defaultLocalPath)
}
baseDir := filepath.Dir(customPath)
_, err := StatDir(baseDir)
if err != nil {
if trace.IsNotFound(err) {
if err := os.MkdirAll(baseDir, teleport.PrivateDirMode); err != nil {
return "", trace.ConvertSystemError(err)
}
} else {
return "", trace.Wrap(err)
}
}
return customPath, nil
}
// IsDir is a helper function to quickly check if a given path is a valid directory
func IsDir(path string) bool {
fi, err := os.Stat(path)
if err == nil {
return fi.IsDir()
}
return false
}
// NormalizePath normalises path, evaluating symlinks and converting local
// paths to absolute
func NormalizePath(path string, evaluateSymlinks bool) (string, error) {
s, err := filepath.Abs(path)
if err != nil {
return "", trace.ConvertSystemError(err)
}
if evaluateSymlinks {
s, err = filepath.EvalSymlinks(s)
if err != nil {
return "", trace.ConvertSystemError(err)
}
}
return s, nil
}
// OpenFileAllowingUnsafeLinks opens a file, if the path includes a symlink, the returned os.File will be resolved to
// the actual file. This will return an error if the file is not found or is a directory.
func OpenFileAllowingUnsafeLinks(path string) (*os.File, error) {
return openFile(path, true /* allowSymlink */, true /* allowMultipleHardlinks */)
}
// OpenFileNoUnsafeLinks opens a file, ensuring it's an actual file and not a directory or symlink. Depending on
// the os, it may also prevent hardlinks. This is important because MacOS allows hardlinks without validating write
// permissions (similar to a symlink in that regard).
func OpenFileNoUnsafeLinks(path string) (*os.File, error) {
return openFile(path, false /* allowSymlink */, runtime.GOOS != "darwin" /* allowMultipleHardlinks */)
}
func openFile(path string, allowSymlink, allowMultipleHardlinks bool) (*os.File, error) {
newPath, err := NormalizePath(path, allowSymlink)
if err != nil {
return nil, trace.Wrap(err)
}
var fi os.FileInfo
if allowSymlink {
fi, err = os.Stat(newPath)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
} else {
components := strings.Split(newPath, string(os.PathSeparator))
var subPath string
for _, p := range components {
subPath = filepath.Join(subPath, p)
if subPath == "" {
subPath = string(os.PathSeparator)
}
fi, err = os.Lstat(subPath)
if err != nil {
return nil, trace.ConvertSystemError(err)
} else if fi.Mode().Type()&os.ModeSymlink != 0 {
return nil, trace.BadParameter("opening file %s, symlink not allowed in path: %s", path, subPath)
}
}
}
if !allowMultipleHardlinks {
// hardlinks can only exist at the end file, not for directories within the path
if linkCount, ok := getHardLinkCount(fi); ok && linkCount > 1 {
return nil, trace.BadParameter("file has hardlink count greater than 1: %s", path)
}
}
if fi.IsDir() {
return nil, trace.BadParameter("%s is not a file", path)
}
f, err := os.Open(newPath)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
return f, nil
}
// StatFile stats path, returns error if it exists but a directory.
func StatFile(path string) (os.FileInfo, error) {
newPath, err := NormalizePath(path, true)
if err != nil {
return nil, trace.Wrap(err)
}
fi, err := os.Stat(newPath)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
if fi.IsDir() {
return nil, trace.BadParameter("%v is not a file", path)
}
return fi, nil
}
// StatDir stats directory, returns error if file exists, but not a directory
func StatDir(path string) (os.FileInfo, error) {
fi, err := os.Stat(path)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
if !fi.IsDir() {
return nil, trace.BadParameter("%v is not a directory", path)
}
return fi, nil
}
// FSTryWriteLock tries to grab write lock, returns ErrUnsuccessfulLockTry
// if lock is already acquired by someone else
func FSTryWriteLock(filePath string) (unlock func() error, err error) {
fileLock := flock.New(getPlatformLockFilePath(filePath))
locked, err := fileLock.TryLock()
if err != nil {
return nil, trace.ConvertSystemError(err)
}
if !locked {
return nil, trace.Retry(ErrUnsuccessfulLockTry, "")
}
return fileLock.Unlock, nil
}
// FSWriteLock tries to grab write lock and block if lock is already acquired by someone else.
func FSWriteLock(filePath string) (unlock func() error, err error) {
fileLock := flock.New(getPlatformLockFilePath(filePath))
if err := fileLock.Lock(); err != nil {
return nil, trace.ConvertSystemError(err)
}
return fileLock.Unlock, nil
}
// FSTryWriteLockTimeout tries to grab write lock, it's doing it until locks is acquired, or timeout is expired,
// or context is expired.
func FSTryWriteLockTimeout(ctx context.Context, filePath string, timeout time.Duration) (unlock func() error, err error) {
fileLock := flock.New(getPlatformLockFilePath(filePath))
timedCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
if _, err := fileLock.TryLockContext(timedCtx, FSLockRetryDelay); err != nil {
return nil, trace.ConvertSystemError(err)
}
return fileLock.Unlock, nil
}
// FSTryReadLock tries to grab shared lock, returns ErrUnsuccessfulLockTry
// if lock is already acquired by someone else
func FSTryReadLock(filePath string) (unlock func() error, err error) {
fileLock := flock.New(getPlatformLockFilePath(filePath))
locked, err := fileLock.TryRLock()
if err != nil {
return nil, trace.ConvertSystemError(err)
}
if !locked {
return nil, trace.Retry(ErrUnsuccessfulLockTry, "")
}
return fileLock.Unlock, nil
}
// FSReadLock tries to grab shared lock and block if lock is already acquired by someone else.
func FSReadLock(filePath string) (unlock func() error, err error) {
fileLock := flock.New(getPlatformLockFilePath(filePath))
if err := fileLock.RLock(); err != nil {
return nil, trace.ConvertSystemError(err)
}
return fileLock.Unlock, nil
}
// FSTryReadLockTimeout tries to grab read lock, it's doing it until locks is acquired, or timeout is expired,
// or context is expired.
func FSTryReadLockTimeout(ctx context.Context, filePath string, timeout time.Duration) (unlock func() error, err error) {
fileLock := flock.New(getPlatformLockFilePath(filePath))
timedCtx, cancel := context.WithTimeout(ctx, timeout)
defer cancel()
if _, err := fileLock.TryRLockContext(timedCtx, FSLockRetryDelay); err != nil {
return nil, trace.ConvertSystemError(err)
}
return fileLock.Unlock, nil
}
// RemoveAllSecure is similar to [os.RemoveAll] but leverages [RemoveSecure] to delete files so that they are
// overwritten. This helps guard against hardware attacks on magnetic disks.
func RemoveAllSecure(path string) error {
if path == "" {
// match behavior from os.RemoveAll
return nil
}
// Match os.RemoveAll protections in not permitting removal of "." directories
// This check comes directly from https://cs.opensource.google/go/go/+/refs/tags/go1.21.1:src/os/removeall_at.go;l=24
if path == "." || (len(path) >= 2 && path[len(path)-1] == '.' && os.IsPathSeparator(path[len(path)-2])) {
return &os.PathError{Op: "RemoveAllSecure", Path: path, Err: syscall.EINVAL} // error type matches os.RemoveAll
}
info, err := os.Lstat(path)
switch {
case err != nil && os.IsNotExist(err):
return nil
case err != nil:
return trace.ConvertSystemError(err)
case !info.IsDir():
return removeSecure(path, info)
}
var removeErrors []error
files, err := os.ReadDir(path)
if err != nil {
// Don't fail fast, allow removal at end to be attempted.
removeErrors = append(removeErrors, err)
}
// It's possible for a partial file list to be returned even if an error above was returned.
for _, f := range files {
if err := RemoveAllSecure(filepath.Join(path, f.Name())); err != nil {
removeErrors = append(removeErrors, err)
}
}
if err := os.Remove(path); err != nil {
removeErrors = append(removeErrors, err)
}
switch len(removeErrors) {
case 1:
return trace.ConvertSystemError(removeErrors[0])
case 0:
return nil
default:
return trace.NewAggregate(removeErrors...)
}
}
// RemoveSecure attempts to securely delete the file by first overwriting the file with random data three times
// followed by calling os.Remove(filePath).
func RemoveSecure(filePath string) error {
info, err := os.Lstat(filePath)
if err != nil && os.IsNotExist(err) {
return err
}
// Don't fast return on other errors, still allow removeSecure to attempt removal.
return removeSecure(filePath, info)
}
func removeSecure(filePath string, fi os.FileInfo) error {
if fi.Mode().Type()&os.ModeSymlink != 0 {
return os.Remove(filePath)
}
f, openErr := os.OpenFile(filePath, os.O_WRONLY, 0)
switch {
case os.IsNotExist(openErr):
return trace.ConvertSystemError(openErr)
case openErr != nil:
// Attempt delete anyway.
return trace.ConvertSystemError(os.Remove(filePath))
}
defer f.Close()
if runtime.GOOS == "windows" {
// On windows, os.Remove() will fail if there are any open handles to the
// file, including in other processes. Skip overwrite to avoid leaving
// files in a broken state.
closeErr := trace.ConvertSystemError(f.Close())
removeErr := trace.ConvertSystemError(os.Remove(filePath))
return trace.NewAggregate(closeErr, removeErr)
} else {
removeErr := os.Remove(filePath)
if f != nil {
for range 3 {
if err := overwriteFile(f, fi); err != nil {
break
}
}
}
return trace.ConvertSystemError(removeErr)
}
}
func overwriteFile(f *os.File, fi os.FileInfo) error {
// Rounding up to 4k to hide the original file size. 4k was chosen because it's a common block size.
const block = 4096
size := fi.Size() / block * block
if fi.Size()%block != 0 {
size += block
}
_, copyErr := io.CopyN(f, rand.Reader, size)
// Attempt sync regardless of above error
syncErr := f.Sync() // sync to ensure commit to hardware
if copyErr != nil {
return trace.Wrap(copyErr)
} else if syncErr != nil {
return trace.Wrap(syncErr)
}
return nil
}
// RemoveFileIfExist removes file if exits.
func RemoveFileIfExist(filePath string) error {
if !FileExists(filePath) {
return nil
}
if err := os.Remove(filePath); err != nil {
return trace.ConvertSystemError(err)
}
return nil
}
func RecursiveChown(dir string, uid, gid int) error {
// First, walk the directory to gather a list of files and directories to update before we open up to modifications
var pathsToUpdate []string
err := filepath.WalkDir(dir, func(path string, d fs.DirEntry, err error) error {
if err != nil {
return trace.Wrap(err)
}
pathsToUpdate = append(pathsToUpdate, path)
return nil
})
if err != nil {
return trace.Wrap(err)
}
// filepath.WalkDir is documented to walk the paths in lexical order, iterating
// in the reverse order ensures that files are always Lchowned before their parent directory
for i := len(pathsToUpdate) - 1; i >= 0; i-- {
path := pathsToUpdate[i]
if err := os.Lchown(path, uid, gid); err != nil {
if errors.Is(err, os.ErrNotExist) {
// Unexpected condition where file was removed after discovery.
continue
}
return trace.Wrap(err)
}
}
return nil
}
func CopyFile(src, dest string, perm os.FileMode) (err error) {
srcFile, err := os.Open(src)
if err != nil {
return trace.ConvertSystemError(err)
}
defer srcFile.Close()
destFile, err := os.OpenFile(dest, os.O_RDWR|os.O_CREATE|os.O_TRUNC, perm)
if err != nil {
return trace.ConvertSystemError(err)
}
defer func() {
err = trace.NewAggregate(err, trace.Wrap(destFile.Close()))
}()
_, err = destFile.ReadFrom(srcFile)
if err != nil {
return trace.ConvertSystemError(err)
}
return nil
}
// RecursivelyCopy will copy a directory from src to dest, if the
// directory exists, files will be overwritten. The skip paramater, if
// provided, will be passed the source and destination paths, and will
// skip files upon returning true
func RecursiveCopy(src, dest string, skip func(src, dest string) (bool, error)) error {
return trace.Wrap(fs.WalkDir(os.DirFS(src), ".", func(path string, d fs.DirEntry, err error) error {
if err != nil {
return trace.Wrap(err)
}
absSrcPath := filepath.Join(src, path)
destPath := filepath.Join(dest, path)
info, err := d.Info()
if err != nil {
return trace.Wrap(err)
}
originalPerm := info.Mode().Perm()
if skip != nil {
doSkip, err := skip(absSrcPath, destPath)
if err != nil {
return trace.Wrap(err)
}
if doSkip {
return nil
}
}
if d.IsDir() {
err := os.Mkdir(destPath, originalPerm)
if os.IsExist(err) {
return nil
}
return trace.ConvertSystemError(err)
}
if d.Type().IsRegular() {
if err := CopyFile(absSrcPath, destPath, originalPerm); err != nil {
return trace.Wrap(err)
}
return nil
}
if info.Mode().Type()&os.ModeSymlink != 0 {
linkDest, err := os.Readlink(absSrcPath)
if err != nil {
return trace.ConvertSystemError(err)
}
if err := os.Symlink(linkDest, destPath); err != nil {
return trace.ConvertSystemError(err)
}
return nil
}
return nil
}))
}
// CreateExclusiveFile creates a file only if it does not exist to prevent overwriting
// existing files.
func CreateExclusiveFile(path string, mode os.FileMode) (*os.File, error) {
out, err := os.OpenFile(path, os.O_CREATE|os.O_EXCL|os.O_WRONLY, mode)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
return out, nil
}
//go:build !windows
/*
* 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 utils
import (
"os"
"syscall"
)
// On non-windows we just lock the target file itself.
func getPlatformLockFilePath(path string) string {
return path
}
func getHardLinkCount(fi os.FileInfo) (uint64, bool) {
if statT, ok := fi.Sys().(*syscall.Stat_t); ok {
// we must do a cast here because this will be uint16 on OSX
//nolint:unconvert // the cast is only necessary for macOS
return uint64(statT.Nlink), true
} else {
return 0, false
}
}
/*
* 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 utils
import (
"bytes"
"errors"
"io"
"net/http"
"github.com/gravitational/trace"
)
// GetAndReplaceRequestBody returns the request body and replaces the drained
// body reader with an [io.NopCloser] allowing for further body processing by
// http transport.
// If memory exhaustion is a concern, it is the caller's responsibility to wrap
// the request body in an [io.LimitReader] prior to calling this function.
func GetAndReplaceRequestBody(req *http.Request) ([]byte, error) {
if req.Body == nil || req.Body == http.NoBody {
return []byte{}, nil
}
defer req.Body.Close()
payload, err := io.ReadAll(req.Body)
if err != nil {
return nil, trace.Wrap(err)
}
req.Body = io.NopCloser(bytes.NewReader(payload))
return payload, nil
}
// GetAndReplaceResponseBody returns the response body and replaces the drained
// body reader with [io.NopCloser] allowing for further body processing.
// If memory exhaustion is a concern, it is the caller's responsibility to wrap
// the response body in an [io.LimitReader] prior to calling this function.
func GetAndReplaceResponseBody(response *http.Response) ([]byte, error) {
if response.Body == nil {
return []byte{}, nil
}
defer response.Body.Close()
payload, err := io.ReadAll(response.Body)
if err != nil {
return nil, trace.Wrap(err)
}
response.Body = io.NopCloser(bytes.NewReader(payload))
return payload, nil
}
// ReplaceRequestBody drains the old request body and replaces it with a new one.
func ReplaceRequestBody(req *http.Request, newBody io.ReadCloser) error {
if err := drainAndCloseRequestBody(req); err != nil {
return trace.Wrap(err)
}
req.Body = newBody
return nil
}
// OverwriteRequestBody (over)writes the new data into the given request's body. It attempts to to
// drain close the request body it overwrites.
func OverwriteRequestBody(req *http.Request, data []byte) error {
if err := drainAndCloseRequestBody(req); err != nil {
return trace.Wrap(err)
}
OverwriteRequestBodyNoDrain(req, data)
return nil
}
// OverwriteRequestBodyNoDrain (over)writes the new data into the given request's body. It does not
// close or drain the request body it overwrites.
func OverwriteRequestBodyNoDrain(req *http.Request, data []byte) {
req.Body = io.NopCloser(bytes.NewReader(data))
req.ContentLength = int64(len(data))
}
// RenameHeader moves all values from the old header key to the new header key.
func RenameHeader(header http.Header, oldKey, newKey string) {
if oldKey == newKey {
return
}
for _, value := range header.Values(oldKey) {
header.Add(newKey, value)
}
header.Del(oldKey)
}
// IsRedirect returns true if the status code is a 3xx code.
func IsRedirect(code int) bool {
if code >= http.StatusMultipleChoices && code <= http.StatusPermanentRedirect {
return true
}
return false
}
// GetAnyHeader returns the first non-empty value by the provided keys.
func GetAnyHeader(header http.Header, keys ...string) string {
for _, key := range keys {
if value := header.Get(key); value != "" {
return value
}
}
return ""
}
// GetSingleHeader will return the header value for the key if there is exactly one value present. If the header is
// missing or specified multiple times, an error will be returned.
func GetSingleHeader(headers http.Header, key string) (string, error) {
values := headers.Values(key)
if len(values) > 1 {
return "", trace.BadParameter("multiple %q headers", key)
} else if len(values) == 0 {
return "", trace.NotFound("missing %q headers", key)
} else {
return values[0], nil
}
}
func drainAndCloseRequestBody(req *http.Request) error {
if req.Body != nil {
defer req.Body.Close()
// drain and discard the request body to allow connection reuse.
_, err := io.Copy(io.Discard, req.Body)
if err != nil && !errors.Is(err, io.EOF) {
return trace.Wrap(err)
}
}
return nil
}
// HTTPDoClient is an interface that defines the Do function of http.Client.
type HTTPDoClient interface {
Do(req *http.Request) (*http.Response, error)
}
// HTTPMiddleware defines a HTTP middleware.
type HTTPMiddleware func(next http.Handler) http.Handler
// ChainHTTPMiddlewares wraps an http.Handler with a list of middlewares. Inner
// middlewares should be provided before outer middlewares.
func ChainHTTPMiddlewares(handler http.Handler, middlewares ...HTTPMiddleware) http.Handler {
if len(middlewares) == 0 {
return handler
}
apply := middlewares[0]
middlewares = middlewares[1:]
if apply != nil {
handler = apply(handler)
}
return ChainHTTPMiddlewares(handler, middlewares...)
}
// NoopHTTPMiddleware is a no-operation HTTPMiddleware that returns the
// original handler.
func NoopHTTPMiddleware(next http.Handler) http.Handler {
return next
}
// MaxBytesReader returns an [io.ReadCloser] that wraps an [http.MaxBytesReader]
// to act as a shim for converting from [http.MaxBytesError] to
// [ErrLimitReached].
func MaxBytesReader(w http.ResponseWriter, r io.ReadCloser, n int64) io.ReadCloser {
return &maxBytesReader{ReadCloser: http.MaxBytesReader(w, r, n)}
}
// maxBytesReader wraps an [http.MaxBytesReader] and converts any
// [http.MaxBytesError] to [ErrLimitReached].
type maxBytesReader struct {
io.ReadCloser
}
func (m *maxBytesReader) Read(p []byte) (int, error) {
n, err := m.ReadCloser.Read(p)
// convert [http.MaxBytesError] to our limit error.
var mbErr *http.MaxBytesError
if errors.As(err, &mbErr) {
return n, ErrLimitReached
}
return n, err
}
/*
* 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 utils
import (
"io/fs"
"time"
)
// InMemoryFile stores the required properties to emulate a File in memory
// It contains the File properties like name, size, mode
// It also contains the File contents
// It does not support folders
type InMemoryFile struct {
name string
mode fs.FileMode
modTime time.Time
content []byte
}
func NewInMemoryFile(name string, mode fs.FileMode, modTime time.Time, content []byte) *InMemoryFile {
return &InMemoryFile{
name: name,
mode: mode,
modTime: modTime,
content: content,
}
}
// Name returns the file's name
func (fi *InMemoryFile) Name() string {
return fi.name
}
// Size returns the file size (calculated when writing the file)
func (fi *InMemoryFile) Size() int64 {
return int64(len(fi.content))
}
// Mode returns the fs.FileMode
func (fi *InMemoryFile) Mode() fs.FileMode {
return fi.mode
}
// ModTime returns the last modification time
func (fi *InMemoryFile) ModTime() time.Time {
return fi.modTime
}
// IsDir checks whether the file is a directory
func (fi *InMemoryFile) IsDir() bool {
return false
}
// Sys is platform independent
// InMemoryFile's implementation is no-op
func (fi *InMemoryFile) Sys() any {
return nil
}
// Content returns the file bytes
func (fi *InMemoryFile) Content() []byte {
return fi.content
}
/*
Copyright 2014 The Kubernetes Authors.
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.
*/
package utils
import (
"bytes"
"encoding/json"
"errors"
"io"
"reflect"
"unicode"
"github.com/ghodss/yaml"
"github.com/gravitational/trace"
jsoniter "github.com/json-iterator/go"
kyaml "k8s.io/apimachinery/pkg/util/yaml"
"github.com/gravitational/teleport/api/internalutils/stream"
)
// ToJSON converts a single YAML document into a JSON document
// or returns an error. If the document appears to be JSON the
// YAML decoding path is not used (so that error messages are
// JSON specific).
// Creds to: k8s.io for the code
func ToJSON(data []byte) ([]byte, error) {
if hasJSONPrefix(data) {
return data, nil
}
return yaml.YAMLToJSON(data)
}
var jsonPrefix = []byte("{")
// hasJSONPrefix returns true if the provided buffer appears to start with
// a JSON open brace.
func hasJSONPrefix(buf []byte) bool {
return hasPrefix(buf, jsonPrefix)
}
// Return true if the first non-whitespace bytes in buf is
// prefix.
func hasPrefix(buf []byte, prefix []byte) bool {
trim := bytes.TrimLeftFunc(buf, unicode.IsSpace)
return bytes.HasPrefix(trim, prefix)
}
// FastUnmarshal uses the json-iterator library for fast JSON unmarshalling.
// Note, this function marshals floats with 6 digits precision.
func FastUnmarshal(data []byte, v any) error {
iter := jsoniter.ConfigFastest.BorrowIterator(data)
defer jsoniter.ConfigFastest.ReturnIterator(iter)
iter.ReadVal(v)
if iter.Error != nil {
return trace.Wrap(iter.Error)
}
return nil
}
// SafeConfig uses jsoniter's ConfigFastest settings but enables map key
// sorting to ensure CompareAndSwap checks consistently succeed.
var SafeConfig = jsoniter.Config{
EscapeHTML: false,
MarshalFloatWith6Digits: true, // will lose precision
ObjectFieldMustBeSimpleString: true, // do not unescape object field
SortMapKeys: true,
}.Froze()
// SafeConfigWithIndent is equivalent to SafeConfig except with indentation
// enabled.
var SafeConfigWithIndent = jsoniter.Config{
IndentionStep: 2,
EscapeHTML: false,
MarshalFloatWith6Digits: true, // will lose precision
ObjectFieldMustBeSimpleString: true, // do not unescape object field
SortMapKeys: true,
}.Froze()
// FastMarshal uses the json-iterator library for fast JSON marshaling.
// Note, this function unmarshals floats with 6 digits precision.
func FastMarshal(v any) ([]byte, error) {
data, err := SafeConfig.Marshal(v)
if err != nil {
return nil, trace.Wrap(err)
}
return data, nil
}
// FastMarshal uses the json-iterator library for fast JSON marshaling
// with indentation. Note, this function unmarshals floats with 6 digits precision.
func FastMarshalIndent(v any, prefix, indent string) ([]byte, error) {
data, err := SafeConfig.MarshalIndent(v, prefix, indent)
if err != nil {
return nil, trace.Wrap(err)
}
return data, nil
}
// WriteJSONArray marshals values as a JSON array.
func WriteJSONArray[T any](w io.Writer, values []T) error {
if len(values) == 0 {
values = []T{}
}
return WriteJSON(w, values)
}
// WriteJSONObject marshals m as a JSON object.
func WriteJSONObject[M ~map[K]V, K comparable, V any](w io.Writer, m M) error {
if len(m) == 0 {
_, err := w.Write([]byte("{}"))
return err
}
return WriteJSON(w, m)
}
// WriteJSON marshals multiple documents as a JSON list with indentation.
func WriteJSON(w io.Writer, values any) error {
encoder := json.NewEncoder(w)
encoder.SetIndent("", " ")
err := encoder.Encode(values)
return trace.Wrap(err)
}
// StremJSONArray streams the elements of a stream.Stream as a json array
// with optional indentation (used to stream to CLI).
func StreamJSONArray[T any](items stream.Stream[T], out io.Writer, indent bool) error {
cfg := SafeConfig
if indent {
cfg = SafeConfigWithIndent
}
stream := jsoniter.NewStream(cfg, out, 512)
stream.WriteArrayStart()
var prev bool
for items.Next() {
if prev {
// if a previous item was written to the array, we need to
// write a comma first.
stream.WriteMore()
}
stream.WriteVal(items.Item())
prev = true
}
stream.WriteArrayEnd()
return trace.NewAggregate(items.Done(), stream.Flush())
}
// WriteYAMLArray marshals values as a YAML array.
func WriteYAMLArray[T any](w io.Writer, values []T) error {
if len(values) == 0 {
values = []T{}
}
return writeYAML(w, values)
}
const yamlDocDelimiter = "---"
// WriteYAML detects whether value is a list
// and marshals multiple documents delimited by `---`, otherwise, marshals
// a single value
func WriteYAML(w io.Writer, values any) error {
if reflect.TypeOf(values).Kind() != reflect.Slice {
return trace.Wrap(writeYAML(w, values))
}
// first pass makes sure that all values are documents (objects or maps)
slice := reflect.ValueOf(values)
if slice.Len() == 0 {
_, err := w.Write([]byte("[]"))
return err
}
allDocs := func() bool {
for i := range slice.Len() {
if !isDoc(slice.Index(i)) {
return false
}
}
return true
}
if !allDocs() {
return trace.Wrap(writeYAML(w, values))
}
// second pass can marshal documents
for i := range slice.Len() {
err := writeYAML(w, slice.Index(i).Interface())
if err != nil {
return trace.Wrap(err)
}
if i != slice.Len()-1 {
if _, err := w.Write([]byte(yamlDocDelimiter + "\n")); err != nil {
return trace.Wrap(err)
}
}
}
return nil
}
// isDoc detects whether value constitutes a document
func isDoc(val reflect.Value) bool {
iterations := 0
for val.Kind() == reflect.Interface || val.Kind() == reflect.Pointer {
val = val.Elem()
// preventing cycles
iterations++
if iterations > 10 {
return false
}
}
return val.Kind() == reflect.Struct || val.Kind() == reflect.Map
}
// writeYAML writes marshaled YAML to writer
func writeYAML(w io.Writer, values any) error {
data, err := yaml.Marshal(values)
if err != nil {
return trace.Wrap(err)
}
_, err = w.Write(data)
return trace.Wrap(err)
}
// ReadYAML can unmarshal a stream of documents, used in tests.
func ReadYAML(reader io.Reader) (any, error) {
decoder := kyaml.NewYAMLOrJSONDecoder(reader, 32*1024)
var values []any
for {
var val any
err := decoder.Decode(&val)
if err != nil {
if errors.Is(err, io.EOF) {
if len(values) == 0 {
return nil, trace.BadParameter("no resources found, empty input?")
}
if len(values) == 1 {
return values[0], nil
}
return values, nil
}
return nil, trace.Wrap(err)
}
values = append(values, val)
}
}
/*
* 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 utils
import (
"fmt"
"io"
"os"
"regexp"
"runtime"
"github.com/coreos/go-semver/semver"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/constants"
)
// KernelVersion parses /proc/sys/kernel/osrelease and returns the kernel
// version of the host. This only returns something meaningful on Linux.
func KernelVersion() (*semver.Version, error) {
if runtime.GOOS != constants.LinuxOS {
return nil, trace.BadParameter("requested kernel version on non-Linux host")
}
file, err := OpenFileNoUnsafeLinks("/proc/sys/kernel/osrelease")
if err != nil {
return nil, trace.Wrap(err)
}
defer file.Close()
ver, err := kernelVersion(file)
if err != nil {
return nil, trace.Wrap(err)
}
return ver, nil
}
// kernelVersionRegex extracts the first three digits of a version from
// a kernel version - this strips off any additional digits or additional
// information appended to the kernel version e.g:
// 5.15.68.1-microsoft-standard-WSL2 => 5.15.68
var kernelVersionRegex = regexp.MustCompile(`^\d+\.\d+\.\d+`)
// kernelVersion reads in the kernel version from the reader and returns
// a *semver.Version.
func kernelVersion(reader io.Reader) (*semver.Version, error) {
buf, err := io.ReadAll(reader)
if err != nil {
return nil, trace.Wrap(err)
}
s := kernelVersionRegex.FindString(string(buf))
if s == "" {
return nil, trace.BadParameter(
"unable to extract kernel semver from string %q",
string(buf),
)
}
ver, err := semver.NewVersion(s)
if err != nil {
return nil, trace.Wrap(err)
}
return ver, nil
}
const btfFile = "/sys/kernel/btf/vmlinux"
// HasBTF checks that the kernel has been compiled with BTF support and
// that the type information can be opened. Returns nil if BTF is there
// and accessible, otherwise an error describing the problem.
func HasBTF() error {
if runtime.GOOS != constants.LinuxOS {
return trace.BadParameter("requested kernel version on non-Linux host")
}
file, err := OpenFileNoUnsafeLinks(btfFile)
if err == nil {
file.Close()
return nil
}
if os.IsNotExist(err) {
return fmt.Errorf("%v was not found. Make sure the kernel was compiled with BTF support (CONFIG_DEBUG_INFO_BTF)", btfFile)
}
return fmt.Errorf("failed to open %v: %w", btfFile, err)
}
// Teleport
// Copyright (C) 2026 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 utils
import "sync"
// KeyLock provides per-key mutual exclusion. Only one caller per key can
// proceed at a time, while unrelated keys are processed concurrently.
//
// The zero value is ready to use. A KeyLock must not be copied after first use.
type KeyLock[K comparable] struct {
mu sync.Mutex
m map[K]*keyLockEntry
}
type keyLockEntry struct {
mu sync.Mutex
refCount int
}
// Lock locks the given key. It blocks until the key is available.
func (k *KeyLock[K]) Lock(key K) {
k.mu.Lock()
if k.m == nil {
k.m = make(map[K]*keyLockEntry)
}
entry, exists := k.m[key]
if !exists {
entry = &keyLockEntry{}
k.m[key] = entry
}
entry.refCount++
k.mu.Unlock()
entry.mu.Lock()
}
// Unlock unlocks the given key. It is a runtime error if the key is not
// locked on entry to Unlock.
func (k *KeyLock[K]) Unlock(key K) {
k.mu.Lock()
defer k.mu.Unlock()
entry := k.m[key]
if entry == nil {
panic("keylock: unlock of unlocked key")
}
if entry.refCount == 0 {
panic("keylock: ref count is zero (this is a bug)")
}
entry.mu.Unlock()
entry.refCount--
if entry.refCount == 0 {
delete(k.m, key)
}
}
// Acquire locks the given key and returns an unlock function. The unlock
// function is safe to call multiple times only the first call has any effect.
func (k *KeyLock[K]) Acquire(key K) func() {
k.Lock(key)
return sync.OnceFunc(func() {
k.Unlock(key)
})
}
/*
Copyright (c) 2013 The go-github AUTHORS. All rights reserved.
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are
met:
* Redistributions of source code must retain the above copyright
notice, this list of conditions and the following disclaimer.
* Redistributions in binary form must reproduce the above
copyright notice, this list of conditions and the following disclaimer
in the documentation and/or other materials provided with the
distribution.
* Neither the name of Google Inc. nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS
"AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT
LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR
A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT
OWNER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL,
SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT
LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE,
DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY
THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
(INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*/
/*
* 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 utils
import (
"net/http"
"net/url"
"strings"
)
// WebLinks holds the pagination links parsed out of a request header
// conforming to RFC 8288.
type WebLinks struct {
// NextPage is the next page of pagination links.
NextPage string
// PrevPage is the previous page of pagination links.
PrevPage string
// FirstPage is the first page of pagination links.
FirstPage string
// LastPage is the last page of pagination links.
LastPage string
}
// ParseWebLinks partially implements RFC 8288 parsing, enough to support
// GitHub pagination links. See https://tools.ietf.org/html/rfc8288 for more
// details on Web Linking and https://github.com/google/go-github for the API
// client that this function was original extracted from.
//
// Link headers typically look like:
//
// Link: <https://api.github.com/user/teams?page=2>; rel="next",
// <https://api.github.com/user/teams?page=34>; rel="last"
func ParseWebLinks(response *http.Response) WebLinks {
wls := WebLinks{}
if links, ok := response.Header["Link"]; ok && len(links) > 0 {
for _, lnk := range links {
for link := range strings.SplitSeq(lnk, ",") {
segments := strings.Split(strings.TrimSpace(link), ";")
// link must at least have href and rel
if len(segments) < 2 {
continue
}
// ensure href is properly formatted
if !strings.HasPrefix(segments[0], "<") || !strings.HasSuffix(segments[0], ">") {
continue
}
// try to pull out page parameter
link, err := url.Parse(segments[0][1 : len(segments[0])-1])
if err != nil {
continue
}
for _, segment := range segments[1:] {
switch strings.TrimSpace(segment) {
case `rel="next"`:
wls.NextPage = link.String()
case `rel="prev"`:
wls.PrevPage = link.String()
case `rel="first"`:
wls.FirstPage = link.String()
case `rel="last"`:
wls.LastPage = link.String()
}
}
}
}
}
return wls
}
/*
* 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 utils
import (
"net"
"os"
"github.com/gravitational/trace"
)
// GetListenerFile returns file associated with listener
func GetListenerFile(listener net.Listener) (*os.File, error) {
switch t := listener.(type) {
case *net.TCPListener:
f, err := t.File()
return f, trace.Wrap(err)
case *net.UnixListener:
f, err := t.File()
return f, trace.Wrap(err)
}
return nil, trace.BadParameter("unsupported listener: %T", listener)
}
/*
* 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 utils
import (
"context"
"errors"
"io"
"log/slog"
"math/rand/v2"
"net"
"slices"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// NewLoadBalancer returns new load balancer listening on frontend
// and redirecting requests to backends using round robin algo
func NewLoadBalancer(ctx context.Context, frontend NetAddr, backends ...NetAddr) (*LoadBalancer, error) {
return newLoadBalancer(ctx, frontend, roundRobinPolicy(), backends...)
}
// NewRandomLoadBalancer returns new load balancer listening on frontend
// and redirecting requests to backends randomly.
func NewRandomLoadBalancer(ctx context.Context, frontend NetAddr, backends ...NetAddr) (*LoadBalancer, error) {
return newLoadBalancer(ctx, frontend, randomPolicy(), backends...)
}
// newLoadBalancer returns new load balancer with the given load balance policy.
func newLoadBalancer(ctx context.Context, frontend NetAddr, policy loadBalancerPolicy, backends ...NetAddr) (*LoadBalancer, error) {
if ctx == nil {
return nil, trace.BadParameter("missing parameter context")
}
return &LoadBalancer{
frontend: frontend,
ctx: ctx,
backends: backends,
policy: policy,
logger: slog.With(
teleport.ComponentKey, "loadbalancer",
"frontend_addr", frontend.FullAddress(),
),
connections: make(map[NetAddr]map[int64]net.Conn),
}, nil
}
// loadBalancerPolicy selects which backend to send traffic to.
type loadBalancerPolicy func([]NetAddr) (NetAddr, error)
// roundRobinPolicy selects backends in sequential order
func roundRobinPolicy() loadBalancerPolicy {
next := -1
return func(backends []NetAddr) (NetAddr, error) {
if len(backends) == 0 {
return NetAddr{}, trace.ConnectionProblem(nil, "no backends")
}
next++
if next >= len(backends) {
next = 0
}
return backends[next], nil
}
}
// randomPolicy selects backends in a random order.
func randomPolicy() loadBalancerPolicy {
return func(backends []NetAddr) (NetAddr, error) {
if len(backends) == 0 {
return NetAddr{}, trace.ConnectionProblem(nil, "no backends")
}
i := rand.N(len(backends))
return backends[i], nil
}
}
// LoadBalancer is a simple load balancer implementation.
// It does not do any health checking of backends and is not suitable for production usage.
type LoadBalancer struct {
sync.RWMutex
connID int64
logger *slog.Logger
frontend NetAddr
backends []NetAddr
ctx context.Context
policy loadBalancerPolicy
listener net.Listener
connections map[NetAddr]map[int64]net.Conn
PROXYHeader []byte // optional PROXY header that load balancer will send to the backend on every new connection.
}
// trackeConnection adds connection to the connection tracker
func (l *LoadBalancer) trackConnection(backend NetAddr, conn net.Conn) int64 {
l.Lock()
defer l.Unlock()
l.connID++
tracker, ok := l.connections[backend]
if !ok {
tracker = make(map[int64]net.Conn)
l.connections[backend] = tracker
}
tracker[l.connID] = conn
return l.connID
}
// untrackConnection removes connection from connection tracker
func (l *LoadBalancer) untrackConnection(backend NetAddr, id int64) {
l.Lock()
defer l.Unlock()
tracker, ok := l.connections[backend]
if !ok {
return
}
delete(tracker, id)
}
// dropConnections drops connections associated with backend
func (l *LoadBalancer) dropConnections(backend NetAddr) {
tracker := l.connections[backend]
for _, conn := range tracker {
conn.Close()
}
delete(l.connections, backend)
}
// AddBackend adds backend
func (l *LoadBalancer) AddBackend(b NetAddr) {
l.Lock()
defer l.Unlock()
l.backends = append(l.backends, b)
l.logger.DebugContext(l.ctx, "Backends updated", "backends", l.backends)
}
// RemoveBackend removes backend
func (l *LoadBalancer) RemoveBackend(b NetAddr) error {
l.Lock()
defer l.Unlock()
for i := range l.backends {
if l.backends[i] == b {
l.backends = slices.Delete(l.backends, i, i+1)
l.dropConnections(b)
return nil
}
}
return trace.NotFound("lb has no backend matching: %+v", b)
}
func (l *LoadBalancer) nextBackend() (NetAddr, error) {
l.Lock()
defer l.Unlock()
backend, err := l.policy(l.backends)
if err != nil {
return NetAddr{}, trace.Wrap(err)
}
return backend, nil
}
func (l *LoadBalancer) closeListener() {
l.Lock()
defer l.Unlock()
if l.listener == nil {
return
}
l.listener.Close()
}
func (l *LoadBalancer) Close() error {
l.closeListener()
return nil
}
// Listen creates a listener on the frontend addr
func (l *LoadBalancer) Listen() error {
var err error
l.listener, err = net.Listen(l.frontend.AddrNetwork, l.frontend.Addr)
if err != nil {
return trace.ConvertSystemError(err)
}
l.logger.DebugContext(l.ctx, "created listening socket",
"listen_addr", logutils.StringerAttr(l.listener.Addr()),
)
return nil
}
// Addr returns the frontend listener address. Call this after Listen,
// otherwise Addr returns nil.
func (l *LoadBalancer) Addr() net.Addr {
if l.listener == nil {
return nil
}
return l.listener.Addr()
}
// Serve starts accepting connections
func (l *LoadBalancer) Serve() error {
for {
conn, err := l.listener.Accept()
if err != nil {
if IsUseOfClosedNetworkError(err) {
return trace.Wrap(err, "listener is closed")
}
select {
case <-l.ctx.Done():
return trace.Wrap(net.ErrClosed, "context is closing")
case <-time.After(5. * time.Second):
l.logger.DebugContext(l.ctx, "Backoff on network error")
}
} else {
go l.forwardConnection(conn)
}
}
}
func (l *LoadBalancer) forwardConnection(conn net.Conn) {
err := l.forward(conn)
if err != nil {
l.logger.WarnContext(l.ctx, "Failed to forward connection", "error", err)
}
}
func (l *LoadBalancer) forward(conn net.Conn) error {
defer conn.Close()
backend, err := l.nextBackend()
if err != nil {
return trace.Wrap(err)
}
connID := l.trackConnection(backend, conn)
defer l.untrackConnection(backend, connID)
backendConn, err := net.Dial(backend.AddrNetwork, backend.Addr)
if err != nil {
return trace.ConvertSystemError(err)
}
defer backendConn.Close()
if len(l.PROXYHeader) > 0 {
if _, err := backendConn.Write(l.PROXYHeader); err != nil {
return trace.Wrap(err)
}
}
backendConnID := l.trackConnection(backend, backendConn)
defer l.untrackConnection(backend, backendConnID)
logger := l.logger.With(
"source_addr", logutils.StringerAttr(conn.RemoteAddr()),
"dest_addr", logutils.StringerAttr(backendConn.RemoteAddr()),
)
logger.DebugContext(l.ctx, "forwarding data")
messagesC := make(chan error, 2)
go func() {
defer conn.Close()
defer backendConn.Close()
_, err := io.Copy(conn, backendConn)
messagesC <- err
}()
go func() {
defer conn.Close()
defer backendConn.Close()
_, err := io.Copy(backendConn, conn)
messagesC <- err
}()
var lastErr error
for range 2 {
select {
case err := <-messagesC:
if err != nil && !errors.Is(err, io.EOF) {
logger.WarnContext(l.ctx, "connection problem", "error", err)
lastErr = err
}
case <-l.ctx.Done():
return trace.ConnectionProblem(nil, "context is closing")
}
}
return lastErr
}
/*
* 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 utils
import (
"net"
"strings"
"github.com/gravitational/trace"
)
// ClientIPFromAddr extracts the client IP from a network address.
// For bufconn test addresses it returns the literal string "bufconn".
func ClientIPFromAddr(addr net.Addr) (string, error) {
if addr == nil {
return "", trace.BadParameter("missing client IP")
}
s := addr.String()
// bufconn peers don't include host:port, so use a stable synthetic
// key for request/connection limiting in tests.
if s == "bufconn" && addr.Network() == "bufconn" {
return "bufconn", nil
}
clientIP, _, err := net.SplitHostPort(s)
if err != nil {
return "", trace.BadParameter("missing client IP")
}
return clientIP, nil
}
// ClientIPFromConn extracts host from provided remote address.
func ClientIPFromConn(conn net.Conn) (string, error) {
clientRemoteAddr := conn.RemoteAddr()
clientIP, _, err := net.SplitHostPort(clientRemoteAddr.String())
if err != nil {
return "", trace.Wrap(err)
}
return clientIP, nil
}
// FindMatchingProxyDNS checks if a given request host or app fqdn matches any of the specified proxy DNS names.
// It compares the hostnames without considering the port numbers.
// If a match is found, the method returns the original proxy DNS name (including its port if present).
// If no match is found, it returns the first proxy DNS name from the list.
//
// Parameters:
// - requestHostnameOrFQDN: A string representing the host in the request, which may include a port.
// - proxyDNSNames: A slice of strings representing possible DNS names for a proxy, each of which may include a port.
//
// Returns:
// - A string representing the matching proxy DNS name with its port, or the first proxy DNS name if no matches are found.
func FindMatchingProxyDNS(requestHostnameOrFQDN string, proxyDNSNames []string) string {
if requestHostnameOrFQDN == "" || len(proxyDNSNames) == 0 {
return ""
}
// Remove port from request host if present.
normalizedRequestHost := strings.Split(requestHostnameOrFQDN, ":")[0]
hostParts := strings.Split(normalizedRequestHost, ".")
// Iterate over each possible suffix of requestHostOrFQDN parts
for start := range hostParts {
possibleHost := strings.Join(hostParts[start:], ".")
for _, proxyDNSName := range proxyDNSNames {
// Normalize proxy DNS name by removing port if present
normalizedProxyDNSName := strings.Split(proxyDNSName, ":")[0]
if possibleHost == normalizedProxyDNSName {
return proxyDNSName
}
}
}
// If no match found, return the first proxyDNSName as fallback
return proxyDNSNames[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/>.
*/
// TODO(nklaassen): evaluate the risks and utility of allowing traits to be used
// as regular expressions. The only thing blocking this today is that all trait
// values are lists and the regex must be a single value. It could be possible
// to write:
// `{{regexp.match(email.local(head(external.trait_name)))}}`
package parse
import (
"fmt"
"net/mail"
"regexp"
"slices"
"strings"
"unicode"
"unicode/utf8"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/typical"
)
const (
// EmailLocalFnName is a name for email.local function
EmailLocalFnName = "email.local"
// RegexpMatchFnName is a name for regexp.match function.
RegexpMatchFnName = "regexp.match"
// RegexpNotMatchFnName is a name for regexp.not_match function.
RegexpNotMatchFnName = "regexp.not_match"
// RegexpReplaceFnName is a name for regexp.replace function.
RegexpReplaceFnName = "regexp.replace"
)
// LabelSelectorSpec parses a string like 'name=value,"long name"="quoted value"` into a map like
// { "name" -> "value", "long name" -> "quoted value" }.
func LabelSelectorSpec(spec string) (map[string]string, error) {
tokens, err := tokenizeLabelSpec(spec)
if err != nil {
return nil, err
}
// break tokens in pairs and put into a map:
labels := make(map[string]string)
for i := 0; i < len(tokens); i += 2 {
labels[tokens[i]] = tokens[i+1]
}
return labels, nil
}
// MultiValueLabelSelectorSpec parses a string like 'name=value,name=other,"long name"="quoted value"`
// into a map like { "name" -> ["value", "other"], "long name" -> ["quoted value"] }.
// Similar to LabelSelectorSpec but allows repeated key values stored into a slice.
// Duplicate values for the same key are dropped.
//
// Multi valued labels are supported for role resources e.g. node_labels.
func MultiValueLabelSelectorSpec(spec string) (map[string][]string, error) {
tokens, err := tokenizeLabelSpec(spec)
if err != nil {
return nil, err
}
// break tokens in pairs and put into a map, appending repeated keys:
labels := make(map[string][]string)
for i := 0; i < len(tokens); i += 2 {
key := tokens[i]
val := tokens[i+1]
if !slices.Contains(labels[key], val) {
labels[key] = append(labels[key], val)
}
}
return labels, nil
}
// tokenizeLabelSpec breaks a label spec like 'name=value,"long name"="quoted value"'
// into a list of key/value tokens (name, value, long name, quoted value...).
func tokenizeLabelSpec(spec string) ([]string, error) {
var tokens []string
openQuotes := false
var tokenStart, assignCount int
specLen := len(spec)
// tokenize the label spec:
for i, ch := range spec {
endOfToken := false
// end of line?
if i+utf8.RuneLen(ch) == specLen {
i += utf8.RuneLen(ch)
endOfToken = true
}
switch ch {
case '"':
openQuotes = !openQuotes
case '=', ',', ';':
if !openQuotes {
endOfToken = true
if ch == '=' {
assignCount++
}
}
}
if endOfToken && i > tokenStart {
tokens = append(tokens, strings.TrimSpace(strings.Trim(spec[tokenStart:i], `"`)))
tokenStart = i + 1
}
}
// simple validation of tokenization: must have an even number of tokens (because they're pairs)
// and the number of such pairs must be equal the number of assignments
if len(tokens)%2 != 0 || assignCount != len(tokens)/2 {
return nil, fmt.Errorf("invalid label spec: '%s', should be 'key=value'", spec)
}
return tokens, nil
}
var (
traitsTemplateParser = mustNewTraitsTemplateParser()
matcherParser = mustNewMatcherParser()
reVariable = regexp.MustCompile(
// prefix is anything that is not { or }
`^(?P<prefix>[^}{]*)` +
// variable is anything in brackets {{}} that is not { or }
`{{(?P<expression>\s*[^}{]*\s*)}}` +
// suffix is anything that is not { or }
`(?P<suffix>[^}{]*)$`,
)
)
// TraitsTemplateExpression can interpolate user trait values into a string
// template to produce some values.
type TraitsTemplateExpression struct {
// prefix is a prefix of the expression
prefix string
// suffix is a suffix of the expression
suffix string
// expr is the expression AST
expr traitsTemplateExpression
}
// NewTraitsTemplateExpression parses expressions like {{external.foo}}, {{internal.bar}},
// or {{user.metadata.name}},
// or a literal value like "prod". Call Interpolate on the returned Expression
// to get the final value based on user traits.
func NewTraitsTemplateExpression(value string) (*TraitsTemplateExpression, error) {
match := reVariable.FindStringSubmatch(value)
if len(match) == 0 {
if strings.Contains(value, "{{") || strings.Contains(value, "}}") {
return nil, trace.BadParameter(
"%q is using template brackets '{{' or '}}', however expression does not parse, make sure the format is {{expression}}",
value,
)
}
expr := typical.LiteralExpr[traitsTemplateEnv, []string]{
Value: []string{value},
}
return &TraitsTemplateExpression{expr: expr}, nil
}
prefix, value, suffix := match[1], match[2], match[3]
expr, err := parseTraitsTemplateExpression(value)
if err != nil {
return nil, trace.Wrap(err)
}
return &TraitsTemplateExpression{
prefix: strings.TrimLeftFunc(prefix, unicode.IsSpace),
suffix: strings.TrimRightFunc(suffix, unicode.IsSpace),
expr: expr,
}, nil
}
// Interpolate interpolates the variable adding prefix and suffix if present.
// The returned error is trace.NotFound in case the expression contains a variable
// and this variable is not found on any trait, nil in case of success,
// and BadParameter otherwise.
func (e *TraitsTemplateExpression) Interpolate(varValidation func(namespace, name string) error, traits map[string][]string) ([]string, error) {
return e.InterpolateWithUser(varValidation, "", traits)
}
// InterpolateWithUser interpolates the variable adding prefix and suffix if
// present, with optional Teleport username.
func (e *TraitsTemplateExpression) InterpolateWithUser(varValidation func(namespace, name string) error, username string, traits map[string][]string) ([]string, error) {
result, err := e.expr.Evaluate(traitsTemplateEnv{
username: username,
traits: traits,
traitValidator: varValidation,
})
if err != nil {
return nil, trace.Wrap(err)
}
var out []string
for _, val := range result {
// Filter out values that mapped to the empty string.
if len(val) > 0 {
out = append(out, e.prefix+val+e.suffix)
}
}
return out, nil
}
type traitsTemplateEnv struct {
username string
traits map[string][]string
traitValidator func(namespace, name string) error
}
type traitsTemplateExpression typical.Expression[traitsTemplateEnv, []string]
func parseTraitsTemplateExpression(exprString string) (traitsTemplateExpression, error) {
expr, err := traitsTemplateParser.Parse(exprString)
return expr, trace.Wrap(err)
}
func mustNewTraitsTemplateParser() *typical.CachedParser[traitsTemplateEnv, []string] {
parser, err := newTraitsTemplateParser()
if err != nil {
panic(trace.Wrap(err, "failed to create template parser (this is a bug)"))
}
return parser
}
func newTraitsTemplateParser() (*typical.CachedParser[traitsTemplateEnv, []string], error) {
traitsVariable := func(name string) typical.Variable {
return typical.DynamicMapFunction(func(e traitsTemplateEnv, key string) ([]string, error) {
if e.traitValidator != nil {
if err := e.traitValidator(name, key); err != nil {
return nil, trace.Wrap(err)
}
}
values, ok := e.traits[key]
if !ok {
return nil, trace.NotFound("trait not found: %s.%s", name, key)
}
return values, nil
})
}
parser, err := typical.NewCachedParser[traitsTemplateEnv, []string](typical.ParserSpec[traitsTemplateEnv]{
Variables: map[string]typical.Variable{
"external": traitsVariable("external"),
"internal": traitsVariable("internal"),
"user.metadata.name": typical.DynamicVariable(func(e traitsTemplateEnv) ([]string, error) {
if e.username == "" {
return nil, trace.NotFound("user.metadata.name is not available in this context")
}
return []string{e.username}, nil
}),
},
Functions: map[string]typical.Function{
EmailLocalFnName: typical.UnaryFunction[traitsTemplateEnv](EmailLocal),
RegexpReplaceFnName: typical.TernaryFunction[traitsTemplateEnv](RegexpReplace),
},
}, typical.WithInvalidNamespaceHack())
return parser, trace.Wrap(err)
}
// EmailLocal returns a new list which is a result of getting the local part of
// each email from the input list.
func EmailLocal(inputs []string) ([]string, error) {
return stringListMap(inputs, func(email string) (string, error) {
if email == "" {
return "", trace.BadParameter(
"found empty %q argument",
EmailLocalFnName,
)
}
addr, err := mail.ParseAddress(email)
if err != nil {
return "", trace.BadParameter(
"failed to parse %q argument %q: %s",
EmailLocalFnName,
email,
err,
)
}
parts := strings.SplitN(addr.Address, "@", 2)
if len(parts) != 2 {
return "", trace.BadParameter(
"could not find local part in %q argument %q, %q",
EmailLocalFnName,
email,
addr.Address,
)
}
return parts[0], nil
})
}
// RegexpReplace returns a new list which is the result of replacing each instance
// of [match] with [replacement] for each item in the input list.
func RegexpReplace(inputs []string, match string, replacement string) ([]string, error) {
re, err := newRegexp(match, false)
if err != nil {
return nil, trace.Wrap(err, "invalid regexp %q", match)
}
return stringListMap(inputs, func(in string) (string, error) {
// Filter out inputs which do not match the regexp at all.
if !re.MatchString(in) {
return "", nil
}
return re.ReplaceAllString(in, replacement), nil
})
}
// MatchExpression is a match expression.
type MatchExpression struct {
// prefix is a prefix of the expression
prefix string
// suffix is a suffix of the expression
suffix string
// matcher is the matcher in the expression
matcher Matcher
}
// Matcher matches strings against some internal criteria (e.g. a regexp)
type Matcher interface {
Match(in string) bool
}
// MatcherFn converts function to a matcher interface
type MatcherFn func(in string) bool
// Match matches string against a regexp
func (fn MatcherFn) Match(in string) bool {
return fn(in)
}
// NewAnyMatcher returns a matcher function based
// on incoming values
func NewAnyMatcher(in []string) (Matcher, error) {
matchers := make([]Matcher, len(in))
for i, v := range in {
m, err := NewMatcher(v)
if err != nil {
return nil, trace.Wrap(err)
}
matchers[i] = m
}
return MatcherFn(func(in string) bool {
for _, m := range matchers {
if m.Match(in) {
return true
}
}
return false
}), nil
}
// NewMatcher parses a matcher expression. Currently supported expressions:
// - string literal: `foo`
// - wildcard expression: `*` or `foo*bar`
// - regexp expression: `^foo$`
// - regexp function calls:
// - positive match: `{{regexp.match("foo.*")}}`
// - negative match: `{{regexp.not_match("foo.*")}}`
//
// These expressions do not support variable interpolation (e.g.
// `{{internal.logins}}`), like Expression does.
func NewMatcher(value string) (*MatchExpression, error) {
match := reVariable.FindStringSubmatch(value)
if len(match) == 0 {
if strings.Contains(value, "{{") || strings.Contains(value, "}}") {
return nil, trace.BadParameter(
"%q is using template brackets '{{' or '}}', however expression does not parse, make sure the format is {{expression}}",
value,
)
}
re, err := newRegexp(value, true)
if err != nil {
return nil, trace.Wrap(err, "parsing match expression")
}
return &MatchExpression{
matcher: matcher{re},
}, nil
}
prefix, value, suffix := match[1], match[2], match[3]
matcher, err := parseMatcherExpression(value)
if err != nil {
return nil, trace.Wrap(err)
}
return &MatchExpression{
prefix: prefix,
suffix: suffix,
matcher: matcher,
}, nil
}
func (e *MatchExpression) Match(in string) bool {
if !strings.HasPrefix(in, e.prefix) || !strings.HasSuffix(in, e.suffix) {
return false
}
in = strings.TrimPrefix(in, e.prefix)
in = strings.TrimSuffix(in, e.suffix)
return e.matcher.Match(in)
}
// match expressions currently have no environment (you can't access any traits
// or other variables).
type matcherEnv struct{}
func parseMatcherExpression(raw string) (Matcher, error) {
matchExpr, err := matcherParser.Parse(raw)
if err != nil {
return nil, trace.Wrap(err, "parsing match expression")
}
matcher, err := matchExpr.Evaluate(matcherEnv{})
return matcher, trace.Wrap(err, "evaluating match expression")
}
func mustNewMatcherParser() *typical.CachedParser[matcherEnv, Matcher] {
parser, err := newMatcherParser()
if err != nil {
panic(trace.Wrap(err, "failed to create match parser (this is a bug)"))
}
return parser
}
func newMatcherParser() (*typical.CachedParser[matcherEnv, Matcher], error) {
parser, err := typical.NewCachedParser[matcherEnv, Matcher](typical.ParserSpec[matcherEnv]{
Functions: map[string]typical.Function{
RegexpMatchFnName: typical.UnaryFunction[matcherEnv](regexpMatch),
RegexpNotMatchFnName: typical.UnaryFunction[matcherEnv](regexpNotMatch),
},
})
return parser, trace.Wrap(err)
}
func regexpMatch(match string) (Matcher, error) {
re, err := newRegexp(match, false)
if err != nil {
return nil, trace.Wrap(err, "parsing argument to regexp.match")
}
return matcher{re}, nil
}
func regexpNotMatch(match string) (Matcher, error) {
re, err := newRegexp(match, false)
if err != nil {
return nil, trace.Wrap(err, "parsing argument to regexp.not_match")
}
return notMatcher{re}, nil
}
type matcher struct {
re *regexp.Regexp
}
func (m matcher) Match(in string) bool {
return m.re.MatchString(in)
}
type notMatcher struct {
re *regexp.Regexp
}
func (n notMatcher) Match(in string) bool {
return !n.re.MatchString(in)
}
func stringListMap(inputs []string, f func(string) (string, error)) ([]string, error) {
out := make([]string, 0, len(inputs))
for _, input := range inputs {
mapped, err := f(input)
if err != nil {
return nil, trace.Wrap(err)
}
// Filter out values that mapped to the empty string.
if len(mapped) == 0 {
continue
}
out = append(out, mapped)
}
return out, nil
}
func newRegexp(raw string, escape bool) (*regexp.Regexp, error) {
if escape {
if !strings.HasPrefix(raw, "^") || !strings.HasSuffix(raw, "$") {
// replace glob-style wildcards with regexp wildcards
// for plain strings, and quote all characters that could
// be interpreted in regular expression
raw = "^" + utils.GlobToRegexp(raw) + "$"
}
}
re, err := regexp.Compile(raw)
if err != nil {
return nil, trace.BadParameter(
"failed to parse regexp %q: %v",
raw,
err,
)
}
return re, nil
}
/*
* 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 utils
import (
"io"
"net"
"sync"
"time"
)
// PipeNetConn implements net.Conn from a provided io.Reader,io.Writer and
// io.Closer
type PipeNetConn struct {
// Locks writing and closing the connection. If both writer & closer refer
// to the same underlying object, simultaneous write and close operations
// introduce a data race (*especially* if that object is a
// `x/crypto/ssh.channel`), so we will use this mutex to serialize write
// and close operations.
mu sync.Mutex
reader io.Reader
writer io.Writer
closer io.Closer
localAddr net.Addr
remoteAddr net.Addr
}
// NewPipeNetConn constructs a new PipeNetConn, providing a net.Conn
// implementation synthesized from the supplied io.Reader, io.Writer &
// io.Closer.
func NewPipeNetConn(reader io.Reader,
writer io.Writer,
closer io.Closer,
fakelocalAddr net.Addr,
fakeRemoteAddr net.Addr) *PipeNetConn {
return &PipeNetConn{
reader: reader,
writer: writer,
closer: closer,
localAddr: fakelocalAddr,
remoteAddr: fakeRemoteAddr,
}
}
func (nc *PipeNetConn) Read(buf []byte) (n int, e error) {
return nc.reader.Read(buf)
}
func (nc *PipeNetConn) Write(buf []byte) (n int, e error) {
nc.mu.Lock()
defer nc.mu.Unlock()
return nc.writer.Write(buf)
}
func (nc *PipeNetConn) Close() error {
nc.mu.Lock()
defer nc.mu.Unlock()
if nc.closer != nil {
return nc.closer.Close()
}
return nil
}
func (nc *PipeNetConn) LocalAddr() net.Addr {
return nc.localAddr
}
func (nc *PipeNetConn) RemoteAddr() net.Addr {
return nc.remoteAddr
}
func (nc *PipeNetConn) SetDeadline(t time.Time) error {
return nil
}
func (nc *PipeNetConn) SetReadDeadline(t time.Time) error {
return nil
}
func (nc *PipeNetConn) SetWriteDeadline(t time.Time) error {
return nil
}
//go:build unix
/*
* 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 utils
import (
"net"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/session/uds"
)
// DualPipeNetConn creates a pipe to connect a client and a server. The
// two net.Conn instances are wrapped in an PipeNetConn which holds the source and
// destination addresses.
//
// The pipe is constructed from a syscall.Socketpair instead of a net.Pipe because
// the synchronous nature of net.Pipe causes it to deadlock when attempting to perform
// TLS or SSH handshakes.
func DualPipeNetConn(srcAddr net.Addr, dstAddr net.Addr) (net.Conn, net.Conn, error) {
client, server, err := uds.NewSocketpair(uds.SocketTypeStream)
if err != nil {
return nil, nil, trace.Wrap(err)
}
serverConn := NewConnWithAddr(server, dstAddr, srcAddr)
clientConn := NewConnWithAddr(client, srcAddr, dstAddr)
return serverConn, clientConn, nil
}
/*
* 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 utils
import (
"context"
"io"
"os"
"github.com/gravitational/trace"
)
// CombinedStdio reads from standard input and writes to standard output.
// Closing a CombinedStdio does nothing, successfully.
type CombinedStdio struct{}
// Read reads from [os.Stdin].
func (CombinedStdio) Read(p []byte) (int, error) {
return os.Stdin.Read(p)
}
// Write writes to [os.Stdout].
func (CombinedStdio) Write(p []byte) (int, error) {
return os.Stdout.Write(p)
}
// ReadFrom copies data from [os.Stdout] to the provided [io.Reader].
func (CombinedStdio) ReadFrom(r io.Reader) (n int64, err error) {
return os.Stdout.ReadFrom(r)
}
// WriteTo copies data from [os.Stdin] to the provided [io.Writer].
func (CombinedStdio) WriteTo(w io.Writer) (n int64, err error) {
return os.Stdin.WriteTo(w)
}
func (CombinedStdio) Close() error {
return nil
}
// ProxyConn launches a double-copy loop that proxies traffic between the
// provided client and server connections.
//
// Exits when one or both copies stop, or when the context is canceled, and
// closes both connections.
func ProxyConn(ctx context.Context, client, server io.ReadWriteCloser) error {
errCh := make(chan error, 2)
defer server.Close()
defer client.Close()
go func() {
defer server.Close()
defer client.Close()
_, err := io.Copy(server, client)
errCh <- err
}()
go func() {
defer server.Close()
defer client.Close()
_, err := io.Copy(client, server)
errCh <- err
}()
var errors []error
for range 2 {
select {
case err := <-errCh:
if err != nil && !IsOKNetworkError(err) {
errors = append(errors, err)
}
case <-ctx.Done():
// Cause(ctx) returns ctx.Err() if no cause is provided.
return trace.Wrap(context.Cause(ctx))
}
}
return trace.NewAggregate(errors...)
}
/*
* 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 utils
import (
"regexp"
"strings"
"github.com/gravitational/trace"
)
var reProxyJump = regexp.MustCompile(
// optional username, note that outside group
`(?:(?P<username>[^\:]+)@)?(?P<hostport>[^\@]+)`,
)
// maxProxyJumpLen is the maximum accepted length of a proxy jump string.
const maxProxyJumpLen = 4096
// JumpHost is a target jump host
type JumpHost struct {
// Username to login as
Username string
// Addr is a target addr
Addr NetAddr
}
// ParseProxyJump parses strings like user@host:port,bob@host:port
func ParseProxyJump(in string) ([]JumpHost, error) {
if in == "" {
return nil, trace.BadParameter("missing proxyjump")
}
if len(in) > maxProxyJumpLen {
return nil, trace.BadParameter("proxyjump too long: %d bytes (max %d)", len(in), maxProxyJumpLen)
}
parts := strings.Split(in, ",")
out := make([]JumpHost, 0, len(parts))
for _, part := range parts {
match := reProxyJump.FindStringSubmatch(strings.TrimSpace(part))
if len(match) == 0 {
return nil, trace.BadParameter("could not parse %q, expected format user@host:port,user@host:port", in)
}
addr, err := ParseAddr(match[2])
if err != nil {
return nil, trace.Wrap(err)
}
out = append(out, JumpHost{Username: match[1], Addr: *addr})
}
return out, nil
}
/*
* 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 utils
import (
"crypto/rand"
"encoding/hex"
"iter"
"math/big"
mathrand "math/rand/v2"
"time"
"github.com/gravitational/trace"
)
// CryptoRandomHex returns a hex-encoded random string generated
// with a crypto-strong pseudo-random generator. The length parameter
// controls how many random bytes are generated, and the returned
// hex string will be twice the length. An error is returned when
// fewer bytes were generated than length.
func CryptoRandomHex(length int) (string, error) {
randomBytes := make([]byte, length)
if _, err := rand.Read(randomBytes); err != nil {
return "", trace.Wrap(err)
}
return hex.EncodeToString(randomBytes), nil
}
// RandomDuration returns a duration in a range [0, max)
func RandomDuration(max time.Duration) time.Duration {
randomVal, err := rand.Int(rand.Reader, big.NewInt(int64(max)))
if err != nil {
return max / 2
}
return time.Duration(randomVal.Int64())
}
// ShuffleVisit yields the items of a slice in random order, while arranging
// them in the same order at the head of the slice. Exiting early from the
// iterator will result in a slice that's partially shuffled - specifically,
// pulling N items from the iterator will also arrange for the slice to contain
// the same items in the same order at indices 0 through N-1. The slice is
// updated as items are yielded from the iterator, so the first N items are
// fixed in position and inspectable during the iteration at step N.
func ShuffleVisit[S ~[]E, E any](s S) iter.Seq2[int, E] {
return func(yield func(int, E) bool) {
for i := range len(s) {
j := mathrand.N(len(s))
// swapping here (instead of swapping after the yield) ensures that
// pulling items from the iterator also puts them in order at the
// beginning of the slice, otherwise there would be a difference
// between exhausting the iterator and exiting early; the items are
// also accessible during the iteration
s[0], s[j] = s[j], s[0]
if !yield(i, s[0]) {
return
}
s = s[1:]
}
}
}
/*
* 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 utils
import (
"maps"
"regexp"
"slices"
"strings"
"github.com/gravitational/trace"
lru "github.com/hashicorp/golang-lru/v2"
"github.com/gravitational/teleport/api/types"
)
// ContainsExpansion returns true if value contains
// expansion syntax, e.g. $1 or ${10}
func ContainsExpansion(val string) bool {
return reExpansion.FindStringIndex(val) != nil
}
// GlobToRegexp replaces glob-style standalone wildcard values
// with real .* regexp-friendly values, does not modify regexp-compatible values,
// quotes non-wildcard values
func GlobToRegexp(in string) string {
return replaceWildcard.ReplaceAllString(regexp.QuoteMeta(in), "(.*)")
}
// ErrReplaceRegexNotFound is a marker error returned by
// [ReplaceRegexp], [RegexpWithConfig], and [ReplaceRegexpWith] to
// indicate no matches were found.
var ErrReplaceRegexNotFound = &trace.NotFoundError{Message: "no match found"}
// ReplaceRegexp replaces value in string, accepts regular expression and simplified
// wildcard syntax, it has several important differences with standard lib
// regexp replacer:
// * Wildcard globs '*' are treated as regular expression .* expression
// * Expression is treated as regular expression if it starts with ^ and ends with $
// * Full match is expected, partial replacements ignored
// * If there is no match, returns [ErrReplaceRegexNotFound]
func ReplaceRegexp(expression string, replaceWith string, input string) (string, error) {
expr, err := RegexpWithConfig(expression, RegexpConfig{})
if err != nil {
return "", trace.Wrap(err)
}
return ReplaceRegexpWith(expr, replaceWith, input)
}
type regexKey struct {
expression string
ignoreCase bool
}
// regexpCache interns compiled regular expressions to improve performance.
var regexpCache = mustCache[regexKey, *regexp.Regexp](2000)
func replaceRegexCached(expression string, config RegexpConfig) (*regexp.Regexp, error) {
key := regexKey{expression: expression, ignoreCase: config.IgnoreCase}
if expr, ok := regexpCache.Get(key); ok {
return expr, nil
}
expression = expressionToRegexp(expression)
if config.IgnoreCase {
expression = "(?i)" + expression
}
expr, err := regexp.Compile(expression)
if err != nil {
return nil, trace.BadParameter("%s", err)
}
regexpCache.Add(key, expr)
return expr, nil
}
// RegexpWithConfig compiles a regular expression given some configuration.
// There are several important differences with standard lib (see ReplaceRegexp).
func RegexpWithConfig(expression string, config RegexpConfig) (*regexp.Regexp, error) {
expr, err := replaceRegexCached(expression, config)
return expr, trace.Wrap(err)
}
// ReplaceRegexpWith replaces string in a given regexp.
func ReplaceRegexpWith(expr *regexp.Regexp, replaceWith string, input string) (string, error) {
index := expr.FindStringIndex(input)
if index == nil {
// The returned error is intentionally not wrapped to avoid
// capturing stack traces. This method is used by authorization
// logic and the additional overhead of strack trace capturing
// is a performance bottleneck.
return "", ErrReplaceRegexNotFound
}
return expr.ReplaceAllString(input, replaceWith), nil
}
// RegexpConfig defines the configuration of the regular expression matcher
type RegexpConfig struct {
// IgnoreCase specifies whether matching is case-insensitive
IgnoreCase bool
}
// KubeResourceMatchesRegex checks whether the input matches any of the given
// expressions.
// This function returns as soon as it finds the first match or when MatchString
// returns an error.
// This function supports regex expressions in the Name and Namespace fields,
// but not for the Kind field.
// The wildcard (*) expansion is also supported.
func KubeResourceMatchesRegexWithVerbsCollector(input types.KubernetesResource, resources []types.KubernetesResource) (bool, []string, error) {
verbs := map[string]struct{}{}
matchedAny := false
for _, resource := range resources {
if input.Kind != resource.Kind && resource.Kind != types.Wildcard {
continue
}
if ok, err := MatchString(input.APIGroup, resource.APIGroup); err != nil {
return false, nil, trace.Wrap(err)
} else if !ok {
continue
}
if ok, err := MatchString(input.Name, resource.Name); err != nil {
return false, nil, trace.Wrap(err)
} else if !ok {
continue
}
if ok, err := MatchString(input.Namespace, resource.Namespace); err != nil {
return false, nil, trace.Wrap(err)
} else if !ok {
continue
}
matchedAny = true
if slices.Contains(resource.Verbs, types.Wildcard) {
return true, []string{types.Wildcard}, nil
}
for _, verb := range resource.Verbs {
verbs[verb] = struct{}{}
}
}
return matchedAny, slices.Collect(maps.Keys(verbs)), nil
}
// KubeResourceMatchesRegex checks whether the input matches any of the given
// expressions.
// This function returns as soon as it finds the first match or when matchString
// returns an error.
// This function supports regex expressions in the Name and Namespace fields,
// but not for the Kind field.
// The wildcard (*) expansion is also supported.
// input is the resource we are checking for access.
// resources is a list of resources that the user has access to - collected from
// their roles that match the Kubernetes cluster where the resource is defined.
// cond is the deny or allow condition of the role that we are evaluating.
func KubeResourceMatchesRegex(input types.KubernetesResource, isClusterWideResource bool, resources []types.KubernetesResource, cond types.RoleConditionType) (bool, error) {
if len(input.Verbs) != 1 {
return false, trace.BadParameter("only one verb is supported, input: %v", input.Verbs)
}
verb := input.Verbs[0]
// If the user is list/read/watch a namespace, they should be able to see the
// namespace they have resources defined for.
// This is a special case because we don't want to require the user to have
// access to the namespace resource itself.
// This is only allowed for the list/read/watch verbs because we don't want
// to allow the user to create/update/delete a namespace they don't have
// permissions for.
targetsReadOnlyNamespace := input.Kind == "namespaces" &&
slices.Contains([]string{types.KubeVerbGet, types.KubeVerbList, types.KubeVerbWatch}, verb)
for _, resource := range resources {
// If the resource has a wildcard verb, it matches all verbs.
// Otherwise, the resource must have the verb we're looking for otherwise
// it doesn't match.
// When the resource has a wildcard verb, we only allow one verb in the
// resource input.
if !IsVerbAllowed(resource.Verbs, verb) {
continue
}
switch {
case targetsReadOnlyNamespace && cond == types.Allow && resource.Kind != "namespaces" && resource.Namespace != "":
// If the user requests a read-only namespace get/list/watch, they should
// be able to see the list of namespaces they have resources defined in.
// This means that if the user has access to pods in the "foo" namespace,
// they should be able to see the "foo" namespace in the list of namespaces
// but only if the request is read-only.
if ok, err := MatchString(input.Name, resource.Namespace); err != nil || ok {
return ok, trace.Wrap(err)
}
case targetsReadOnlyNamespace && cond == types.Allow && resource.Kind == "namespaces" && resource.Name != "":
if ok, err := MatchString(input.Name, resource.Name); err != nil || ok {
return ok, trace.Wrap(err)
}
case input.Kind == "namespaces":
if input.Kind != resource.Kind && resource.Kind != types.Wildcard {
continue
}
if ok, err := MatchString(input.APIGroup, resource.APIGroup); err != nil {
return false, trace.Wrap(err)
} else if !ok {
continue
}
targetNamespace := resource.Namespace
if resource.Kind == "namespaces" {
targetNamespace = resource.Name
} else if resource.Kind == types.Wildcard && (resource.Namespace == "" || resource.Namespace == types.Wildcard) {
targetNamespace = resource.Name
}
if ok, err := MatchString(input.Name, targetNamespace); err != nil || ok {
return ok, trace.Wrap(err)
}
// No match.
continue
default:
if input.Kind != resource.Kind && resource.Kind != types.Wildcard {
continue
}
if ok, err := MatchString(input.APIGroup, resource.APIGroup); err != nil {
return false, trace.Wrap(err)
} else if !ok {
continue
}
if ok, err := MatchString(input.Name, resource.Name); err != nil {
return false, trace.Wrap(err)
} else if !ok {
continue
}
if input.Namespace == "" && resource.Namespace != "" && resource.Namespace != types.Wildcard {
continue
}
// At this point everything else matched. If we match the namespace as well, we have a match.
if ok, err := MatchString(input.Namespace, resource.Namespace); err != nil || ok {
return ok, trace.Wrap(err)
}
}
}
return false, nil
}
// KubeResourceCouldMatchRules assess whether the user is permitted to perform its request
// based on the defined kubernetes_resource rules. The aim is to catch cases when the user
// has no access and present then a more user-friendly error message instead of returning
// an empty list.
// This function is not responsible for enforcing access rules.
func KubeResourceCouldMatchRules(input types.KubernetesResource, isClusterWideResource bool, resources []types.KubernetesResource, cond types.RoleConditionType) (bool, error) {
if len(input.Verbs) != 1 {
return false, trace.BadParameter("only one verb is supported, input: %v", input.Verbs)
}
if input.Name != "" {
return false, trace.BadParameter("name is not supported for KubeResourceCouldMatchRules")
}
verb := input.Verbs[0]
isDeny := cond == types.Deny
// If the user is allowed to list/read/watch a resource, they should be able to see the
// namespace in which the resource is.
// This is a special case because we don't want to require the user to have
// access to the namespace resource itself.
// This is only allowed for the list/read/watch verbs because we don't want
// to allow the user to create/update/delete a namespace they don't have
// permissions for.
targetsReadOnlyNamespace := input.Kind == "namespaces" &&
slices.Contains([]string{types.KubeVerbGet, types.KubeVerbList, types.KubeVerbWatch}, verb)
for _, resource := range resources {
// If the resource has a wildcard verb, it matches all verbs.
// Otherwise, the resource must have the verb we're looking for otherwise
// it doesn't match.
// When the resource has a wildcard verb, we only allow one verb in the
// resource input.
if !IsVerbAllowed(resource.Verbs, verb) {
continue
}
switch {
case targetsReadOnlyNamespace && isDeny:
// For read-only namespace request, match the deny only if there is an explicit deny,
// i.e., if we have a wildcard deny, we should still be able to get namespaces.
// If the group doesn't match and is not wildcard, skip.
if resource.Kind != "namespaces" {
continue // The only possible way to match in deny is to have an explicit 'namespaces' rule.
}
if ok, err := MatchString(input.Name, resource.Name); err != nil || ok {
return ok, trace.Wrap(err)
}
continue
case targetsReadOnlyNamespace && !isDeny && resource.Kind != "namespaces" && resource.Namespace != "":
// If the user requests a read-only namespace get/list/watch, they should
// be able to see the list of namespaces they have resources defined in.
// This means that if the user has access to pods in the "foo" namespace,
// they should be able to see the "foo" namespace in the list of namespaces
// but only if the request is read-only.
return true, nil
default:
// If the kind doesn't match and is not wildcard, skip.
if input.Kind != resource.Kind && resource.Kind != types.Wildcard {
continue
}
// If the group doesn't match and is not wildcard, skip.
if ok, err := MatchString(input.APIGroup, resource.APIGroup); err != nil {
return false, trace.Wrap(err)
} else if !ok {
continue
}
// if the resource is cluster-wide, the command is deny and it's a wildcard resource
// match all resources.
if isClusterWideResource && isDeny && resource.Name == types.Wildcard {
return true, nil
} else if isClusterWideResource {
return !isDeny, nil
}
// If we are listing a namespaced resource, we can't match against a cluster-wide entry.
if isDeny && resource.Namespace == "" {
return false, nil
}
// at this point, the resource is namespaced and if the namespace is empty,
// the user is requesting resources in all namespaces.
// Since he has some rule defined, we should return.
isAllowOrFullDeny := !isDeny || isDeny && resource.Name == types.Wildcard && resource.Namespace == types.Wildcard
if input.Namespace == "" && isAllowOrFullDeny {
return isAllowOrFullDeny, nil
}
if ok, err := MatchString(input.Namespace, resource.Namespace); err != nil {
return false, trace.Wrap(err)
} else if !ok {
continue
}
if !isDeny || isDeny && resource.Name == types.Wildcard {
return !isDeny || isDeny && resource.Name == types.Wildcard, nil
}
}
}
return false, nil
}
// IsVerbAllowed returns true if the verb is allowed in the resource.
// A wildcard verb anywhere in the list matches all verbs,
// otherwise the verb must appear in the list explicitly.
func IsVerbAllowed(allowedVerbs []string, verb string) bool {
return len(allowedVerbs) != 0 && (slices.Contains(allowedVerbs, types.Wildcard) || slices.Contains(allowedVerbs, verb))
}
// SliceMatchesRegex checks if input matches any of the expressions. The
// match is always evaluated as a regex either an exact match or regexp.
func SliceMatchesRegex(input string, expressions []string) (bool, error) {
for _, expression := range expressions {
result, err := MatchString(input, expression)
if err != nil || result {
return result, trace.Wrap(err)
}
}
return false, nil
}
// RegexMatchesAny returns true if [expression] matches any element of
// [inputs]. [expression] support globbing ("env-*") or normal regexp syntax if
// surrounded with ^$ ("^env-.*$").
func RegexMatchesAny(inputs []string, expression string) (bool, error) {
expr, err := compileRegexCached(expression)
if err != nil {
return false, trace.Wrap(err)
}
if slices.ContainsFunc(inputs, expr.MatchString) {
return true, nil
}
return false, nil
}
// mustCache initializes a new [lru.Cache] with the provided size.
// A panic will be triggered if the creation of the cache fails.
func mustCache[K comparable, V any](size int) *lru.Cache[K, V] {
cache, err := lru.New[K, V](size)
if err != nil {
panic(err)
}
return cache
}
// MatchString will match an input against the given expression. The expression is cached for later use.
func MatchString(input, expression string) (bool, error) {
expr, err := compileRegexCached(expression)
if err != nil {
return false, trace.BadParameter("%s", err)
}
// Since the expression is always surrounded by ^ and $ this is an exact
// match for either a plain string (for example ^hello$) or for a regexp
// (for example ^hel*o$).
return expr.MatchString(input), nil
}
// IsRegexp returns true if the expression is a raw regex pattern (starts with ^ and ends with $).
func IsRegexp(expression string) bool {
return strings.HasPrefix(expression, "^") && strings.HasSuffix(expression, "$")
}
// expressionToRegexp converts a Teleport expression to a regexp string.
func expressionToRegexp(expression string) string {
if IsRegexp(expression) {
return expression
}
// replace glob-style wildcards with regexp wildcards
// for plain strings, and quote all characters that could
// be interpreted in regular expression
return "^" + GlobToRegexp(expression) + "$"
}
// CompileExpression compiles the given regex expression with Teleport's custom globbing
// and quoting logic.
func CompileExpression(expression string) (*regexp.Regexp, error) {
expression = expressionToRegexp(expression)
expr, err := regexp.Compile(expression)
if err != nil {
return nil, trace.BadParameter("%s", err)
}
return expr, nil
}
func compileRegexCached(expression string) (*regexp.Regexp, error) {
key := regexKey{expression: expression}
if expr, ok := regexpCache.Get(key); ok {
return expr, nil
}
expr, err := CompileExpression(expression)
if err != nil {
return nil, trace.Wrap(err)
}
regexpCache.Add(key, expr)
return expr, nil
}
var (
replaceWildcard = regexp.MustCompile(`(\\\*)`)
reExpansion = regexp.MustCompile(`\$[^\$]+`)
)
/*
* 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 utils
import (
"time"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/api/utils/retryutils"
)
// HalfJitter is [retryutils.HalfJitter].
//
// Deprecated: use retryutils.HalfJitter.
func HalfJitter(d time.Duration) time.Duration { return retryutils.HalfJitter(d) }
// FullJitter is [retryutils.FullJitter].
//
// Deprecated: use retryutils.FullJitter.
func FullJitter(d time.Duration) time.Duration { return retryutils.FullJitter(d) }
// NewDefaultLinear creates a linear retry with reasonable default parameters for
// attempting to restart "critical but potentially load-inducing" operations, such
// as watcher or control stream resume. Exact parameters are subject to change,
// but this retry will always be configured for automatic reset.
func NewDefaultLinear(clock clockwork.Clock) *retryutils.Linear {
retry, err := retryutils.NewLinear(retryutils.LinearConfig{
First: retryutils.FullJitter(time.Second * 10),
Step: time.Second * 15,
Max: time.Second * 90,
Jitter: retryutils.HalfJitter,
AutoReset: 5,
Clock: clock,
})
if err != nil {
panic("default linear retry misconfigured (this is a bug)")
}
return retry
}
// Copyright 2009 The Go Authors. All rights reserved.
// Use of this source code is governed by a BSD-style
// license that can be found in the LICENSE file.
package utils
import (
"math"
)
const (
uvone = 0x3FF0000000000000
mask = 0x7FF
shift = 64 - 11 - 1
bias = 1023
signMask = 1 << 63
fracMask = 1<<shift - 1
)
// Round returns the nearest integer, rounding half away from zero.
//
// Special cases are:
//
// Round(±0) = ±0
// Round(±Inf) = ±Inf
// Round(NaN) = NaN
//
// Note: Copied from Go standard library to support Go 1.9.7 releases. This
// function was added in the standard library in Go 1.10.
func Round(x float64) float64 {
// Round is a faster implementation of:
//
// func Round(x float64) float64 {
// t := Trunc(x)
// if Abs(x-t) >= 0.5 {
// return t + Copysign(1, x)
// }
// return t
// }
bits := math.Float64bits(x)
e := uint(bits>>shift) & mask
if e < bias {
// Round abs(x) < 1 including denormals.
bits &= signMask // +-0
if e == bias-1 {
bits |= uvone // +-1
}
} else if e < bias+shift {
// Round any abs(x) >= 1 containing a fractional component [0,1).
//
// Numbers with larger exponents are returned unchanged since they
// must be either an integer, infinity, or NaN.
const half = 1 << (shift - 1)
e -= bias
bits += half >> e
bits &^= fracMask >> e
}
return math.Float64frombits(bits)
}
/*
* 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 utils
import "sync/atomic"
// RoundRobin is a helper for distributing load across multiple resources in a round-robin
// fashion.
type RoundRobin[T any] struct {
ct atomic.Uint64
items []T
}
// NewRoundRobin creates a new round-robin inst
func NewRoundRobin[T any](items []T) *RoundRobin[T] {
return &RoundRobin[T]{
items: items,
}
}
// Next gets the next item that is up for use.
func (r *RoundRobin[T]) Next() T {
n := r.ct.Add(1) - 1
l := uint64(len(r.items))
return r.items[int(n%l)]
}
// ForEach applies the supplied closure to each item.
func (r *RoundRobin[T]) ForEach(fn func(T)) {
for _, item := range r.items {
fn(item)
}
}
/*
* 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 utils
import (
"bytes"
"sync"
)
// SlicePool manages a pool of slices
// in attempts to manage memory in go more efficiently
// and avoid frequent allocations
type SlicePool interface {
// Zero zeroes slice
Zero(b []byte)
// Get returns a new or already allocated slice
Get() []byte
// Put returns slice back to the pool
Put(b []byte)
// Size returns a slice size
Size() int64
}
// NewSliceSyncPool returns a new slice pool, using sync.Pool
// of pre-allocated or newly allocated slices of the predefined size and capacity
func NewSliceSyncPool(sliceSize int64) *SliceSyncPool {
s := &SliceSyncPool{
sliceSize: sliceSize,
zeroSlice: make([]byte, sliceSize),
}
s.New = func() any {
slice := make([]byte, s.sliceSize)
return &slice
}
return s
}
// SliceSyncPool is a sync pool of slices (usually large)
// of the same size to optimize memory usage, see sync.Pool for more details
type SliceSyncPool struct {
sync.Pool
sliceSize int64
zeroSlice []byte
}
// Zero zeroes slice of any length
func (s *SliceSyncPool) Zero(b []byte) {
if len(b) <= len(s.zeroSlice) {
// zero all bytes in the slice to avoid
// data lingering in memory
copy(b, s.zeroSlice[:len(b)])
} else {
// use working, but less optimal implementation
for i := range b {
b[i] = 0
}
}
}
// Get returns a new or already allocated slice
func (s *SliceSyncPool) Get() []byte {
pslice := s.Pool.Get().(*[]byte)
return *pslice
}
// Put returns slice back to the pool
func (s *SliceSyncPool) Put(b []byte) {
s.Zero(b)
s.Pool.Put(&b)
}
// Size returns a slice size
func (s *SliceSyncPool) Size() int64 {
return s.sliceSize
}
// NewBufferSyncPool returns a new instance of sync pool of bytes.Buffers
// that creates new buffers with preallocated underlying buffer of size
func NewBufferSyncPool(size int64) *BufferSyncPool {
return &BufferSyncPool{
size: size,
Pool: sync.Pool{
New: func() any {
return bytes.NewBuffer(make([]byte, size))
},
},
}
}
// BufferSyncPool is a sync pool of bytes.Buffer
type BufferSyncPool struct {
sync.Pool
size int64
}
// Put resets the buffer (does not free the memory)
// and returns it back to the pool. Users should be careful
// not to use the buffer (e.g. via Bytes) after it was returned
func (b *BufferSyncPool) Put(buf *bytes.Buffer) {
buf.Reset()
b.Pool.Put(buf)
}
// Get returns a new or already allocated buffer
func (b *BufferSyncPool) Get() *bytes.Buffer {
return b.Pool.Get().(*bytes.Buffer)
}
// Size returns default allocated buffer size
func (b *BufferSyncPool) Size() int64 {
return b.size
}
// FromSlice converts the provided slice to a map using the key function
// to determine the appropriate key per entry. If any duplicates
// exist in the slice, then the entry with the lowest index is used.
func FromSlice[T any](r []T, key func(T) string) map[string]T {
out := make(map[string]T, len(r))
// there may be duplicate resources in the input list.
// by iterating from end to start, the first resource of given name wins.
for i := len(r) - 1; i >= 0; i-- {
res := r[i]
out[key(res)] = res
}
return out
}
/*
* 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 utils
import (
"crypto/sha256"
"crypto/subtle"
"crypto/x509"
"encoding/hex"
"strings"
"github.com/gravitational/trace"
)
// CalculateSPKI the hash value of the SPKI header in a certificate.
func CalculateSPKI(cert *x509.Certificate) string {
sum := sha256.Sum256(cert.RawSubjectPublicKeyInfo)
return "sha256:" + hex.EncodeToString(sum[:])
}
// CheckSPKI the passed in pin against the calculated value from a certificate.
func CheckSPKI(pins []string, certs []*x509.Certificate) error {
// check pins
for _, pin := range pins {
// Check that the format of the pin is valid.
parts := strings.Split(pin, ":")
if len(parts) != 2 {
return trace.BadParameter("invalid format for certificate pin, expected algorithm:pin")
}
if parts[0] != "sha256" {
return trace.BadParameter("sha256 only supported hashing algorithm for certificate pin")
}
}
// Timing of this check depends only on the number of pins and certs, not
// their contents.
outer:
for _, cert := range certs {
for _, pin := range pins {
// Check that that pin itself matches that value calculated from the passed
// in certificate.
if subtle.ConstantTimeCompare([]byte(CalculateSPKI(cert)), []byte(pin)) == 1 {
continue outer
}
}
return trace.BadParameter("%s", errorMessage)
}
return nil
}
var errorMessage = "cluster pin does not match any provided certificate authority pin. " +
"This could have occurred if the Certificate Authority (CA) for the cluster " +
"was rotated, invalidating the old pin. This could also occur if a new HSM was " +
"added. Run \"tctl status\" to compare the pin used to join the cluster to the " +
"actual pin(s) for the cluster."
/*
* Teleport
* Copyright (C) 2026 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 utils
import "strings"
// TrimNonEmpty trims leading and trailing whitespace from s and reports whether
// the trimmed value is non-empty. It is shaped to be a drop-in callback for
// filter-and-transform helpers such as lib/utils/slices.FilterMapUnique, where
// the boolean return decides whether the transformed value is kept.
//
// Typical use is normalizing user-edited string lists (matcher selectors,
// config arrays) where stray whitespace and empty entries should both be
// dropped before downstream processing.
func TrimNonEmpty(s string) (string, bool) {
s = strings.TrimSpace(s)
return s, s != ""
}
/*
* 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 utils
import (
"maps"
"sync"
)
// SyncMap is a generics version of a sync.Map.
type SyncMap[K comparable, V any] struct {
values map[K]V
mu sync.RWMutex
}
// Load returns the value stored in the map for a key.
func (s *SyncMap[K, V]) Load(key K) (value V, ok bool) {
s.mu.RLock()
defer s.mu.RUnlock()
if s.values == nil {
return value, false
}
value, ok = s.values[key]
return value, ok
}
// Store sets the value for a key.
func (s *SyncMap[K, V]) Store(key K, value V) {
s.mu.Lock()
defer s.mu.Unlock()
if s.values == nil {
s.values = make(map[K]V)
}
s.values[key] = value
}
// Delete deletes the value for a key.
func (s *SyncMap[K, V]) Delete(key K) {
s.mu.Lock()
defer s.mu.Unlock()
if s.values == nil {
return
}
delete(s.values, key)
}
// LoadAndDelete loads the value for a key and deletes it if it exists.
func (s *SyncMap[K, V]) LoadAndDelete(key K) (value V, ok bool) {
s.mu.Lock()
defer s.mu.Unlock()
if s.values == nil {
return value, false
}
value, ok = s.values[key]
if ok {
delete(s.values, key)
}
return value, ok
}
// Range calls a function sequentially for each key and value in the map.
// Caution: The map is ony locked while creating a copy of the map values to
// iterate over; it is *not* locked during the actual iteration over that copy,
// nor while f is being evaluated.
func (s *SyncMap[K, V]) Range(f func(key K, value V) bool) {
items := s.Clone()
for key, value := range items {
if !f(key, value) {
return
}
}
}
// Clear clears the underlying map
func (s *SyncMap[K, V]) Clear() {
s.mu.Lock()
defer s.mu.Unlock()
if s.values != nil {
clear(s.values)
}
}
// Set sets the underlying managed map to the supplied value. Note that directly reading
// from or writing to `m` outside of the SyncMap after calling `Set()` may result in a
// data race .
func (s *SyncMap[K, V]) Set(m map[K]V) {
s.mu.Lock()
defer s.mu.Unlock()
s.values = m
}
// Clone creates an un-synchronized shallow clone of the protected map
func (s *SyncMap[K, V]) Clone() map[K]V {
s.mu.RLock()
defer s.mu.RUnlock()
return maps.Clone(s.values)
}
// Len fetches the number of items in the map
func (s *SyncMap[K, V]) Len() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.values)
}
// Read acquires the map read lock and applies the supplied function the
// underlying map, automatically releasing the lock when done. Prefer `Load()`
// when you want to read a single map value. Only prefer `Read()` when
// - you need a coherent view of the map across multiple reads, or
// - ensuring that reference-type map values (e.g. maps, struct on the heap)
// are not modified while being read
func (s *SyncMap[K, V]) Read(fn func(map[K]V)) {
s.mu.RLock()
defer s.mu.RUnlock()
fn(s.values)
}
// Write acquires the map lock and applies the supplied function the
// underlying map, automatically releasing the lock when done.
func (s *SyncMap[K, V]) Write(fn func(map[K]V)) {
s.mu.Lock()
defer s.mu.Unlock()
if s.values == nil {
s.values = make(map[K]V)
}
fn(s.values)
}
/*
* 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 utils
import (
"io"
"sync"
)
type SyncWriter struct {
io.Writer
sync.Mutex
}
func NewSyncWriter(w io.Writer) *SyncWriter {
return &SyncWriter{
Writer: w,
}
}
func (sw *SyncWriter) Write(b []byte) (int, error) {
sw.Lock()
defer sw.Unlock()
return sw.Writer.Write(b)
}
/*
* 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 utils
import (
"time"
"github.com/jonboulle/clockwork"
"google.golang.org/protobuf/types/known/timestamppb"
)
// TimeFromProto converts a protobuf Timestamp to a Go time.Time, preserving
// the zero value across the conversion boundary (standard go/proto timestamp
// conversion doesn't preserve "zeroness").
func TimeFromProto(t *timestamppb.Timestamp) time.Time {
// use the zero time to represent the nil timestamp. note that this is conceptually distinct
// from using t.GetSeconds() == 0 && t.GetNanos() == 0. a timstampb that happens to be created
// targeting the unix epoch isn't necessarily equivalent to a zero go timestamp, since the zero
// value for the go timestamp isn't the unix epoch.
if t == nil {
return time.Time{}
}
return t.AsTime()
}
// TimeIntoProto converts a Go time.Time to a protobuf Timestamp, preserving
// the zero value across the conversion boundary (standard go/proto timestamp
// conversion doesn't preserve "zeroness").
func TimeIntoProto(t time.Time) *timestamppb.Timestamp {
if t.IsZero() {
return nil
}
return timestamppb.New(t)
}
// MinTTL selects the smallest non-zero duration from a and b.
func MinTTL(a, b time.Duration) time.Duration {
if a == 0 {
return b
}
if b == 0 {
return a
}
if a < b {
return a
}
return b
}
// ToTTL converts expiration time to TTL duration
// relative to current time as provided by clock
func ToTTL(c clockwork.Clock, tm time.Time) time.Duration {
now := c.Now().UTC()
if tm.IsZero() || tm.Before(now) {
return 0
}
return tm.Sub(now)
}
/*
* 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 utils
import (
"net"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
)
// ObeyIdleTimeout wraps an existing network connection, closing it if data
// isn't read often enough. The connection will be closed even if Read is never
// called, or if it's called on the underlying connection instead of the
// returned one.
func ObeyIdleTimeout(conn net.Conn, timeout time.Duration) net.Conn {
return obeyIdleTimeoutClock(conn, timeout, clockwork.NewRealClock())
}
// obeyIdleTimeoutClock is [ObeyIdleTimeout] but lets the caller specify an
// arbitrary [clockwork.Clock] to be used for the timer.
func obeyIdleTimeoutClock(conn net.Conn, timeout time.Duration, clock clockwork.Clock) net.Conn {
return &timeoutConn{
Conn: conn,
timeout: timeout,
watchdog: clock.AfterFunc(timeout, func() {
conn.Close()
}),
}
}
func ObeyIdleTimeoutDisarmed(conn net.Conn) (net.Conn, func(time.Duration)) {
return obeyIdleTimeoutDisarmed(conn, clockwork.NewRealClock())
}
func obeyIdleTimeoutDisarmed(conn net.Conn, clock clockwork.Clock) (net.Conn, func(time.Duration)) {
tc := &timeoutConn{
Conn: conn,
}
return tc, func(timeout time.Duration) {
tc.arm(timeout, clock)
}
}
type timeoutConn struct {
net.Conn
timeout time.Duration
mu sync.Mutex
watchdog clockwork.Timer
closed bool
}
func (c *timeoutConn) arm(timeout time.Duration, clock clockwork.Clock) {
c.mu.Lock()
defer c.mu.Unlock()
if !c.closed && c.watchdog == nil {
c.timeout = timeout
c.watchdog = clock.AfterFunc(timeout, func() {
c.Conn.Close()
})
}
}
func (c *timeoutConn) pet() {
c.mu.Lock()
defer c.mu.Unlock()
// if the timer has already fired the underlying net.Conn has been closed or
// will be closed shortly anyway
if c.watchdog != nil && c.watchdog.Stop() {
c.watchdog.Reset(c.timeout)
}
}
// NetConn returns the underlying [net.Conn].
func (c *timeoutConn) NetConn() net.Conn {
return c.Conn
}
// Close implements [io.Closer] and [net.Conn] by closing the underlying
// connection and then stopping the watchdog, if it's still running.
func (c *timeoutConn) Close() error {
err := c.Conn.Close()
c.mu.Lock()
defer c.mu.Unlock()
if c.watchdog != nil {
c.watchdog.Stop()
}
c.closed = true
return trace.Wrap(err)
}
// Read implements [io.Reader] and [net.Conn], petting the watchdog timer if any
// data is successfully read.
func (c *timeoutConn) Read(p []byte) (n int, err error) {
n, err = c.Conn.Read(p)
if n > 0 {
c.pet()
}
// avoid trace.Wrap to maintain the exact errors from the underlying
// connection (like io.EOF)
return n, err
}
/*
* 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 utils
import (
"context"
"crypto/tls"
"crypto/x509"
"net"
"time"
"github.com/gravitational/trace"
)
// TLSConfig returns default TLS configuration strong defaults.
func TLSConfig(cipherSuites []uint16) *tls.Config {
config := &tls.Config{}
SetupTLSConfig(config, cipherSuites)
return config
}
// SetupTLSConfig sets up cipher suites in existing TLS config
func SetupTLSConfig(config *tls.Config, cipherSuites []uint16) {
// If ciphers suites were passed in, use them. Otherwise use the
// Go defaults.
if len(cipherSuites) > 0 {
config.CipherSuites = cipherSuites
}
// pre-v17 Teleport uses a client ticket cache, which doesn't play well with
// verification (both client- and server-side) when using dynamic
// credentials and CAs (in v17+ Teleport)
config.SessionTicketsDisabled = true
config.MinVersion = tls.VersionTLS12
}
// CipherSuiteMapping transforms Teleport formatted cipher suites strings
// into uint16 IDs.
func CipherSuiteMapping(cipherSuites []string) ([]uint16, error) {
out := make([]uint16, 0, len(cipherSuites))
for _, cs := range cipherSuites {
c, ok := cipherSuiteMapping[cs]
if !ok {
return nil, trace.BadParameter("cipher suite not supported: %v", cs)
}
out = append(out, c)
}
return out, nil
}
// VerifyConnectionWithRoots returns a [tls.Config.VerifyConnection] function
// that uses the provided function to generate a [*x509.CertPool] that's used as
// the source of root CAs for verification. Use of this function requires a
// modicum of care: the [*tls.Config] using the returned callback should be
// generated as close to its point of use as possible. An example use for this
// would be something like:
//
// c := utils.TLSConfig(cfg.cipherSuites)
// c.GetClientCertificate = func(cri *tls.CertificateRequestInfo) (*tls.Certificate, error) {
// return cfg.getCert()
// }
// c.ServerName = apiutils.EncodeClusterName(cfg.clusterName)
// c.InsecureSkipVerify = true
// c.VerifyConnection = VerifyConnectionWithRoots(cfg.getRoots)
// httpTransport.TLSClientConfig = c
// clientConn := grpc.NewClient(target, grpc.WithTransportCredentials(credentials.NewTLS(c)))
//
// The necessity of using InsecureSkipVerify is the reason why this construction
// is deliberately not packaged into a utility function, as the stakes must be
// clear to whoever is interacting with the constructed [*tls.Config]. The
// recommended approach is to push the getter functions as close to the point of
// use as possible.
// TODO(tross): Remove when all references are replaced with [VerifyConnection].
func VerifyConnectionWithRoots(getRoots func() (*x509.CertPool, error)) func(cs tls.ConnectionState) error {
return VerifyConnection(time.Now, getRoots)
}
// VerifyConnection returns a [tls.Config.VerifyConnection] function
// that uses the provided function to generate a [*x509.CertPool] that's used as
// the source of root CAs for verification. Use of this function requires a
// modicum of care: the [*tls.Config] using the returned callback should be
// generated as close to its point of use as possible. An example use for this
// would be something like:
//
// c := utils.TLSConfig(cfg.cipherSuites)
// c.GetClientCertificate = func(cri *tls.CertificateRequestInfo) (*tls.Certificate, error) {
// return cfg.getCert()
// }
// c.ServerName = apiutils.EncodeClusterName(cfg.clusterName)
// c.InsecureSkipVerify = true
// c.VerifyConnection = VerifyConnectionWithRoots(cfg.getRoots)
// httpTransport.TLSClientConfig = c
// clientConn := grpc.NewClient(target, grpc.WithTransportCredentials(credentials.NewTLS(c)))
//
// The necessity of using InsecureSkipVerify is the reason why this construction
// is deliberately not packaged into a utility function, as the stakes must be
// clear to whoever is interacting with the constructed [*tls.Config]. The
// recommended approach is to push the getter functions as close to the point of
// use as possible.
func VerifyConnection(now func() time.Time, getRoots func() (*x509.CertPool, error)) func(cs tls.ConnectionState) error {
return func(cs tls.ConnectionState) error {
if cs.ServerName == "" {
return trace.BadParameter("TLS verification requires a server name")
}
roots, err := getRoots()
if err != nil {
return trace.Wrap(err)
}
opts := x509.VerifyOptions{
Roots: roots,
Intermediates: nil,
DNSName: cs.ServerName,
}
if len(cs.PeerCertificates) > 1 {
opts.Intermediates = x509.NewCertPool()
for _, cert := range cs.PeerCertificates[1:] {
opts.Intermediates.AddCert(cert)
}
}
if now != nil {
opts.CurrentTime = now()
}
if _, err := cs.PeerCertificates[0].Verify(opts); err != nil {
return trace.Wrap(err)
}
return nil
}
}
type (
GetCertificateFunc = func() (*tls.Certificate, error)
GetRootsFunc = func() (*x509.CertPool, error)
)
// TLSConn is a `net.Conn` that implements some of the functions defined by the
// `tls.Conn` struct. This interface can be used where it could receive a
// `tls.Conn` wrapped in another connection. For example, in the ALPN Proxy,
// some TLS Connections can be wrapped with ping protocol.
type TLSConn interface {
net.Conn
// ConnectionState returns basic TLS details about the connection.
// More info at: https://pkg.go.dev/crypto/tls#Conn.ConnectionState
ConnectionState() tls.ConnectionState
// Handshake runs the client or server handshake protocol if it has not yet
// been run.
// More info at: https://pkg.go.dev/crypto/tls#Conn.Handshake
Handshake() error
// HandshakeContext runs the client or server handshake protocol if it has
// not yet been run.
// More info at: https://pkg.go.dev/crypto/tls#Conn.HandshakeContext
HandshakeContext(context.Context) error
}
// cipherSuiteMapping is the mapping between Teleport formatted cipher
// suites strings and uint16 IDs.
var cipherSuiteMapping = map[string]uint16{
"tls-ecdhe-ecdsa-with-aes-128-cbc-sha": tls.TLS_ECDHE_ECDSA_WITH_AES_128_CBC_SHA,
"tls-ecdhe-ecdsa-with-aes-256-cbc-sha": tls.TLS_ECDHE_ECDSA_WITH_AES_256_CBC_SHA,
"tls-ecdhe-rsa-with-aes-128-cbc-sha": tls.TLS_ECDHE_RSA_WITH_AES_128_CBC_SHA,
"tls-ecdhe-rsa-with-aes-256-cbc-sha": tls.TLS_ECDHE_RSA_WITH_AES_256_CBC_SHA,
"tls-ecdhe-rsa-with-aes-128-gcm-sha256": tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256,
"tls-ecdhe-ecdsa-with-aes-128-gcm-sha256": tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256,
"tls-ecdhe-rsa-with-aes-256-gcm-sha384": tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
"tls-ecdhe-ecdsa-with-aes-256-gcm-sha384": tls.TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384,
"tls-ecdhe-rsa-with-chacha20-poly1305": tls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305,
"tls-ecdhe-ecdsa-with-chacha20-poly1305": tls.TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305,
}
const (
// DefaultLRUCapacity is a capacity for LRU session cache
DefaultLRUCapacity = 1024
// DefaultCertTTL sets the TTL of the self-signed certificate (1 year)
DefaultCertTTL = (24 * time.Hour) * 365
)
// DefaultCipherSuites returns the default list of cipher suites that
// Teleport supports. By default Teleport only support modern ciphers
// (Chacha20 and AES GCM) and key exchanges which support perfect forward
// secrecy (ECDHE).
//
// Note that TLS_RSA_WITH_AES_128_GCM_SHA{256,384} have been dropped due to
// being banned by HTTP2 which breaks gRPC clients. For more information see:
// https://tools.ietf.org/html/rfc7540#appendix-A. These two can still be
// manually added if needed.
func DefaultCipherSuites() []uint16 {
return []uint16{
tls.TLS_ECDHE_RSA_WITH_CHACHA20_POLY1305,
tls.TLS_ECDHE_ECDSA_WITH_CHACHA20_POLY1305,
tls.TLS_ECDHE_RSA_WITH_AES_128_GCM_SHA256,
tls.TLS_ECDHE_ECDSA_WITH_AES_128_GCM_SHA256,
tls.TLS_ECDHE_RSA_WITH_AES_256_GCM_SHA384,
tls.TLS_ECDHE_ECDSA_WITH_AES_256_GCM_SHA384,
}
}
// RefreshTLSConfigTickets should be called right before cloning a [tls.Config]
// for a one-off use to not break TLS session resumption, as a workaround for
// https://github.com/golang/go/issues/60506 .
func RefreshTLSConfigTickets(c *tls.Config) {
_, _ = c.DecryptTicket(nil, tls.ConnectionState{})
}
/*
* 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 utils
import (
"github.com/google/uuid"
"github.com/gravitational/teleport/lib/fixtures"
)
// UID provides an interface for generating unique identifiers.
type UID interface {
// New returns a new UUID4.
New() string
}
// realUID is a real UID generator.
type realUID struct{}
// NewRealUID returns a new real UID generator.
func NewRealUID() UID {
return &realUID{}
}
// New generates a new UUID4.
func (u *realUID) New() string {
return uuid.New().String()
}
// fakeUID is a fake UID generator used in tests.
type fakeUID struct{}
// NewFakeUID returns a new fake UID generator used in tests.
func NewFakeUID() UID {
return &fakeUID{}
}
// New returns a fake UUID4.
func (u *fakeUID) New() string {
return fixtures.UUID
}
/*
* 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 utils
import (
"archive/tar"
"context"
"errors"
"io"
"log/slog"
"os"
"path"
"path/filepath"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
)
// Extract extracts the contents of the specified tarball under dir. The
// resulting files and directories are created using the current user context.
// Extract will only unarchive files into dir, and will fail if the tarball
// tries to write files outside of dir.
//
// If any paths are specified, only the specified paths are extracted.
// The destination specified in the first matching path is selected.
func Extract(r io.Reader, dir string, paths ...ExtractPath) error {
tarball := tar.NewReader(r)
for {
header, err := tarball.Next()
if errors.Is(err, io.EOF) {
break
} else if err != nil {
return trace.Wrap(err)
}
dirMode, ok := filterHeader(header, paths)
if !ok {
continue
}
err = sanitizeTarPath(header, dir)
if err != nil {
return trace.Wrap(err)
}
if err := extractFile(tarball, header, dir, dirMode); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// ExtractPath specifies a path to be extracted.
type ExtractPath struct {
// Src path and Dst path within the archive to extract files to.
// Directories in the Src path are not included in the extraction dir.
// For example, given foo/bar/file.txt with Src=foo/bar Dst=baz, baz/file.txt results.
// Trailing slashes are always ignored.
Src, Dst string
// Skip extracting the Src path and ignore Dst.
Skip bool
// DirMode is the file mode for implicit parent directories in Dst.
DirMode os.FileMode
}
// filterHeader modifies the tar header by filtering it through the ExtractPaths.
// filterHeader returns false if the tar header should be skipped.
// If no paths are provided, filterHeader assumes the header should be included, and sets
// the mode for implicit parent directories to teleport.DirMaskSharedGroup.
func filterHeader(hdr *tar.Header, paths []ExtractPath) (dirMode os.FileMode, include bool) {
name := path.Clean(hdr.Name)
for _, p := range paths {
src := path.Clean(p.Src)
switch hdr.Typeflag {
case tar.TypeDir:
// If name is a directory, then
// assume src is a directory prefix, or the directory itself,
// and replace that prefix with dst.
if src != "/" {
src += "/" // ensure HasPrefix does not match partial names
}
if !strings.HasPrefix(name, src) {
continue
}
dst := path.Join(p.Dst, strings.TrimPrefix(name, src))
if dst != "/" {
dst += "/" // tar directory headers end in /
}
hdr.Name = dst
return p.DirMode, !p.Skip
default:
// If name is a file, then
// if src is an exact match to the file name, assume src is a file and write directly to dst,
// otherwise, assume src is a directory prefix, and replace that prefix with dst.
if src == name {
hdr.Name = path.Clean(p.Dst)
return p.DirMode, !p.Skip
}
if src != "/" {
src += "/" // ensure HasPrefix does not match partial names
}
if !strings.HasPrefix(name, src) {
continue
}
hdr.Name = path.Join(p.Dst, strings.TrimPrefix(name, src))
return p.DirMode, !p.Skip
}
}
return teleport.DirMaskSharedGroup, len(paths) == 0
}
// extractFile extracts a single file or directory from tarball into dir.
// Uses header to determine the type of item to create
// Based on https://github.com/mholt/archiver
func extractFile(tarball *tar.Reader, header *tar.Header, dir string, dirMode os.FileMode) error {
switch header.Typeflag {
case tar.TypeDir:
return withDir(filepath.Join(dir, header.Name), dirMode, nil)
case tar.TypeBlock, tar.TypeChar, tar.TypeReg, tar.TypeFifo:
return writeFile(filepath.Join(dir, header.Name), tarball, header.FileInfo().Mode(), dirMode)
case tar.TypeLink:
return writeHardLink(filepath.Join(dir, header.Name), filepath.Join(dir, header.Linkname), dirMode)
case tar.TypeSymlink:
return writeSymbolicLink(filepath.Join(dir, header.Name), header.Linkname, dirMode)
default:
slog.WarnContext(context.Background(), "Unsupported type flag for tarball",
"type_flag", header.Typeflag,
"header", header.Name,
)
}
return nil
}
// sanitizeTarPath checks that the tar header paths resolve to a subdirectory
// path, and don't contain file paths or links that could escape the tar file
// like ../../etc/password.
func sanitizeTarPath(header *tar.Header, dir string) error {
// Sanitize all tar paths resolve to within the destination directory.
destPath := filepath.Join(dir, header.Name)
if !strings.HasPrefix(destPath, filepath.Clean(dir)+string(os.PathSeparator)) {
return trace.BadParameter("%s: illegal file path", header.Name)
}
// Ensure link destinations resolve to within the destination directory.
if header.Linkname != "" {
if filepath.IsAbs(header.Linkname) {
if !strings.HasPrefix(filepath.Clean(header.Linkname), filepath.Clean(dir)+string(os.PathSeparator)) {
return trace.BadParameter("%s: illegal link path", header.Linkname)
}
} else {
// Relative paths are relative to filename after extraction to directory.
linkPath := filepath.Join(dir, filepath.Dir(header.Name), header.Linkname)
if !strings.HasPrefix(linkPath, filepath.Clean(dir)+string(os.PathSeparator)) {
return trace.BadParameter("%s: illegal link path", header.Linkname)
}
}
}
return nil
}
func writeFile(path string, r io.Reader, mode, dirMode os.FileMode) error {
err := withDir(path, dirMode, func() error {
// Create file only if it does not exist to prevent overwriting existing
// files (like session recordings).
out, err := CreateExclusiveFile(path, mode)
if err != nil {
return trace.ConvertSystemError(err)
}
_, err = io.Copy(out, r)
return trace.NewAggregate(err, out.Close())
})
return trace.Wrap(err)
}
func writeSymbolicLink(path, target string, dirMode os.FileMode) error {
err := withDir(path, dirMode, func() error {
err := os.Symlink(target, path)
return trace.ConvertSystemError(err)
})
return trace.Wrap(err)
}
func writeHardLink(path, target string, dirMode os.FileMode) error {
err := withDir(path, dirMode, func() error {
err := os.Link(target, path)
return trace.ConvertSystemError(err)
})
return trace.Wrap(err)
}
func withDir(path string, mode os.FileMode, fn func() error) error {
err := os.MkdirAll(filepath.Dir(path), mode)
if err != nil {
return trace.ConvertSystemError(err)
}
if fn == nil {
return nil
}
err = fn()
return trace.Wrap(err)
}
/*
* 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 utils
import (
"unsafe"
"github.com/gravitational/trace"
)
// UnsafeSliceData is a wrapper around unsafe.SliceData
// which ensures that instead of ever returning
// "a non-nil pointer to an unspecified memory address"
// (see unsafe.SliceData documentation), an error is
// returned instead.
func UnsafeSliceData[T any](slice []T) (*T, error) {
if slice == nil || cap(slice) > 0 {
return unsafe.SliceData(slice), nil
}
return nil, trace.BadParameter("non-nil slice had a capacity of zero")
}
// Teleport
// Copyright (C) 2025 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/>.
//go:build !docs
package utils
import (
"github.com/alecthomas/kingpin/v2"
)
// UpdateAppUsageTemplate is a no-op for regular builds. See usage_docs.go for
// doc generation.
func UpdateAppUsageTemplate(*kingpin.Application) {}
// DocsMode is false in regular builds. It is true when building with -tags docs.
const DocsMode = false
/*
* 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 utils
import (
"context"
"errors"
"io"
"io/fs"
"log/slog"
"math/rand/v2"
"net"
"os"
"path/filepath"
"runtime"
"slices"
"strconv"
"strings"
"sync"
"time"
"unicode"
"github.com/gravitational/trace"
"k8s.io/apimachinery/pkg/util/validation"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/constants"
apiutils "github.com/gravitational/teleport/api/utils"
)
// WriteContextCloser provides close method with context
type WriteContextCloser interface {
Close(ctx context.Context) error
io.Writer
}
// WriteCloserWithContext converts ContextCloser to io.Closer,
// whenever new Close method will be called, the ctx will be passed to it
func WriteCloserWithContext(ctx context.Context, closer WriteContextCloser) io.WriteCloser {
return &closerWithContext{
WriteContextCloser: closer,
ctx: ctx,
}
}
type closerWithContext struct {
WriteContextCloser
ctx context.Context
}
// Close closes all resources and returns the result
func (c *closerWithContext) Close() error {
return c.WriteContextCloser.Close(c.ctx)
}
// NilCloser returns closer if it's not nil
// otherwise returns a nop closer
func NilCloser(r io.Closer) io.Closer {
if r == nil {
return &nilCloser{}
}
return r
}
type nilCloser struct {
}
func (*nilCloser) Close() error {
return nil
}
// assert that CloseFunc implement io.Closer.
var _ io.Closer = (CloseFunc)(nil)
// CloseFunc is a helper used to implement io.Closer on a closure.
type CloseFunc func() error
func (cf CloseFunc) Close() error {
return cf()
}
// NopWriteCloser returns a WriteCloser with a no-op Close method wrapping
// the provided Writer w
func NopWriteCloser(r io.Writer) io.WriteCloser {
return nopWriteCloser{r}
}
type nopWriteCloser struct {
io.Writer
}
func (nopWriteCloser) Close() error { return nil }
// Tracer helps to trace execution of functions
type Tracer struct {
// Started records starting time of the call
Started time.Time
// Description is arbitrary description
Description string
}
// NewTracer returns a new tracer
func NewTracer(description string) *Tracer {
return &Tracer{Started: time.Now().UTC(), Description: description}
}
// Start logs start of the trace
func (t *Tracer) Start() *Tracer {
slog.DebugContext(context.Background(), "Tracer started",
"trace", t.Description)
return t
}
// Stop logs stop of the trace
func (t *Tracer) Stop() *Tracer {
slog.DebugContext(context.Background(), "Tracer completed",
"trace", t.Description,
"duration", time.Since(t.Started),
)
return t
}
// ThisFunction returns calling function name
func ThisFunction() string {
var pc [32]uintptr
runtime.Callers(2, pc[:])
return runtime.FuncForPC(pc[0]).Name()
}
// AsBool converts string to bool, in case of the value is empty
// or unknown, defaults to false
func AsBool(v string) bool {
if v == "" {
return false
}
out, _ := apiutils.ParseBool(v)
return out
}
// ParseAdvertiseAddr validates advertise address,
// makes sure it's not an unreachable or multicast address
// returns address split into host and port, port could be empty
// if not specified
func ParseAdvertiseAddr(advertiseIP string) (string, string, error) {
advertiseIP = strings.TrimSpace(advertiseIP)
host := advertiseIP
port := ""
if len(net.ParseIP(host)) == 0 && strings.Contains(advertiseIP, ":") {
var err error
host, port, err = net.SplitHostPort(advertiseIP)
if err != nil {
return "", "", trace.BadParameter("failed to parse address %q", advertiseIP)
}
if _, err := strconv.Atoi(port); err != nil {
return "", "", trace.BadParameter("bad port %q, expected integer", port)
}
if host == "" {
return "", "", trace.BadParameter("missing host parameter")
}
}
ip := net.ParseIP(host)
if len(ip) != 0 {
if ip.IsUnspecified() || ip.IsMulticast() {
return "", "", trace.BadParameter("unreachable advertise IP: %v", advertiseIP)
}
}
return host, port, nil
}
// DNSName extracts DNS name from host:port string,
// returning an error if the hostname is an IP address.
func DNSName(hostport string) (string, error) {
host, err := Host(hostport)
if err != nil {
return "", trace.Wrap(err)
}
if ip := net.ParseIP(host); len(ip) != 0 {
return "", trace.BadParameter("%v is an IP address", host)
}
return host, nil
}
// Host extracts host from host:port string
func Host(hostname string) (string, error) {
if hostname == "" {
return "", trace.BadParameter("missing parameter hostname")
}
// if this is IPv4 or V6, return as is
if ip := net.ParseIP(hostname); len(ip) != 0 {
return hostname, nil
}
// has no indication of port, return, note that
// it will not break ipv6 as it always has at least one colon
if !strings.Contains(hostname, ":") {
return hostname, nil
}
host, _, err := SplitHostPort(hostname)
if err != nil {
return "", trace.Wrap(err)
}
return host, nil
}
// TryHost is a utility function that extracts host from the host:port pair,
// in case of any error returns the original value.
func TryHost(in string) string {
out, err := Host(in)
if err != nil {
return in
}
return out
}
// SplitHostPort splits host and port and checks that host is not empty
func SplitHostPort(hostname string) (string, string, error) {
host, port, err := net.SplitHostPort(hostname)
if err != nil {
return "", "", trace.Wrap(err)
}
if host == "" {
return "", "", trace.BadParameter("empty hostname")
}
return host, port, nil
}
// HostFQDN consists of host UUID and cluster name joined via '.'
func HostFQDN(hostUUID, clusterName string) string {
return hostUUID + "." + clusterName
}
// IsValidHostname checks if a string represents a valid hostname.
func IsValidHostname(hostname string) bool {
for label := range strings.SplitSeq(hostname, ".") {
if len(validation.IsDNS1035Label(label)) > 0 {
return false
}
}
return true
}
// IsValidUnixUser checks if a string represents a valid
// UNIX username.
func IsValidUnixUser(u string) bool {
// See http://www.unix.com/man-page/linux/8/useradd:
//
// On Debian, the only constraints are that usernames must neither start with a dash ('-')
// nor contain a colon (':') or a whitespace (space: ' ', end of line: '\n', tabulation:
// '\t', etc.). Note that using a slash ('/') may break the default algorithm for the
// definition of the user's home directory.
const maxUsernameLen = 32
if len(u) > maxUsernameLen || len(u) == 0 || u[0] == '-' {
return false
}
if strings.ContainsAny(u, ":/") {
return false
}
for _, r := range u {
if unicode.IsSpace(r) || unicode.IsControl(r) {
return false
}
}
return true
}
// ReadPath reads file contents
func ReadPath(path string) ([]byte, error) {
if path == "" {
return nil, trace.NotFound("empty path")
}
s, err := filepath.Abs(path)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
abs, err := filepath.EvalSymlinks(s)
if err != nil {
if errors.Is(err, fs.ErrPermission) {
//do not convert to system error as this loses the ability to compare that it is a permission error
return nil, err
}
return nil, trace.ConvertSystemError(err)
}
bytes, err := os.ReadFile(abs)
if err != nil {
if errors.Is(err, fs.ErrPermission) {
//do not convert to system error as this loses the ability to compare that it is a permission error
return nil, err
}
return nil, trace.ConvertSystemError(err)
}
return bytes, nil
}
type multiCloser struct {
closers []io.Closer
}
func (mc *multiCloser) Close() error {
for _, closer := range mc.closers {
if err := closer.Close(); err != nil {
return trace.Wrap(err)
}
}
return nil
}
// MultiCloser implements io.Close, it sequentially calls Close() on each object
func MultiCloser(closers ...io.Closer) io.Closer {
return &multiCloser{
closers: closers,
}
}
// IsHandshakeFailedError specifies whether this error indicates
// failed handshake
func IsHandshakeFailedError(err error) bool {
if err == nil {
return false
}
return strings.Contains(trace.Unwrap(err).Error(), "ssh: handshake failed")
}
// IsCertExpiredError specifies whether this error indicates
// expired SSH certificate
func IsCertExpiredError(err error) bool {
if err == nil {
return false
}
return strings.Contains(trace.Unwrap(err).Error(), "ssh: cert has expired")
}
// OpaqueAccessDenied returns a generic [trace.NotFoundError] if [err] is a [trace.NotFoundError] or
// a [trace.AccessDeniedError] so as to avoid leaking the existence of secret resources,
// for other error types it returns the original error.
func OpaqueAccessDenied(err error) error {
if trace.IsNotFound(err) || trace.IsAccessDenied(err) {
return trace.NotFound("not found")
}
return trace.Wrap(err)
}
// PortList is a list of TCP ports.
type PortList struct {
ports []string
sync.Mutex
}
// Pop returns a value from the list, it panics if the value is not there
func (p *PortList) Pop() string {
p.Lock()
defer p.Unlock()
if len(p.ports) == 0 {
panic("list is empty")
}
val := p.ports[len(p.ports)-1]
p.ports = p.ports[:len(p.ports)-1]
return val
}
// PopInt returns a value from the list, it panics if not enough values
// were allocated
func (p *PortList) PopInt() int {
i, err := strconv.Atoi(p.Pop())
if err != nil {
panic(err)
}
return i
}
// PortStartingNumber is a starting port number for tests
const PortStartingNumber = 20000
// GetFreeTCPPorts returns n ports starting from port 20000.
func GetFreeTCPPorts(n int, offset ...int) (PortList, error) {
list := make([]string, 0, n)
start := PortStartingNumber
if len(offset) != 0 {
start = offset[0]
}
for i := start; i < start+n; i++ {
list = append(list, strconv.Itoa(i))
}
return PortList{ports: list}, nil
}
// RemoveFromSlice makes a copy of the slice and removes the passed in values from the copy.
func RemoveFromSlice(slice []string, values ...string) []string {
return slices.DeleteFunc(
slices.Clone(slice),
func(s string) bool {
return slices.Contains(values, s)
},
)
}
// ChooseRandomString returns a random string from the given slice.
func ChooseRandomString(slice []string) string {
switch len(slice) {
case 0:
return ""
case 1:
return slice[0]
default:
return slice[rand.N(len(slice))]
}
}
// CheckCertificateFormatFlag checks if the certificate format is valid.
func CheckCertificateFormatFlag(s string) (string, error) {
switch s {
case constants.CertificateFormatStandard, teleport.CertificateFormatOldSSH, teleport.CertificateFormatUnspecified:
return s, nil
default:
return "", trace.BadParameter("invalid certificate format parameter: %q", s)
}
}
// AddrsFromStrings returns strings list converted to address list
func AddrsFromStrings(s apiutils.Strings, defaultPort int) ([]NetAddr, error) {
addrs := make([]NetAddr, len(s))
for i, val := range s {
addr, err := ParseHostPortAddr(val, defaultPort)
if err != nil {
return nil, trace.Wrap(err)
}
addrs[i] = *addr
}
return addrs, nil
}
// FileExists checks whether a file exists at a given path
func FileExists(fp string) bool {
_, err := os.Stat(fp)
return !errors.Is(err, fs.ErrNotExist)
}
// StoreErrorOf stores the error returned by f within *err.
func StoreErrorOf(f func() error, err *error) {
*err = trace.NewAggregate(*err, f())
}
// LimitReader returns a reader that limits bytes from r, and reports an error
// when limit bytes are read.
func LimitReader(r io.Reader, limit int64) io.Reader {
return &limitedReader{
LimitedReader: &io.LimitedReader{R: r, N: limit},
}
}
// limitedReader wraps an [io.LimitedReader] that limits bytes read, and
// reports an error when the read limit is reached.
type limitedReader struct {
*io.LimitedReader
}
func (l *limitedReader) Read(p []byte) (int, error) {
n, err := l.LimitedReader.Read(p)
if l.LimitedReader.N <= 0 {
return n, ErrLimitReached
}
return n, err
}
// ReadAtMost reads up to limit bytes from r, and reports an error
// when limit bytes are read.
func ReadAtMost(r io.Reader, limit int64) ([]byte, error) {
limitedReader := LimitReader(r, limit)
data, err := io.ReadAll(limitedReader)
return data, err
}
// ErrLimitReached means that the read limit is reached.
//
// TODO(gavin): this should be converted to a 413 StatusRequestEntityTooLarge
// in trace.ErrorToCode instead of 429 StatusTooManyRequests.
var ErrLimitReached = &trace.LimitExceededError{Message: "the read limit is reached"}
const (
// CertTeleportUser specifies teleport user
CertTeleportUser = "x-teleport-user"
// CertTeleportUserCA specifies teleport certificate authority
CertTeleportUserCA = "x-teleport-user-ca"
// CertExtensionRole specifies teleport role
CertExtensionRole = "x-teleport-role"
// CertExtensionAuthority specifies teleport authority's name
// that signed this domain
CertExtensionAuthority = "x-teleport-authority"
// CertTeleportClusterName is a name of the teleport cluster
CertTeleportClusterName = "x-teleport-cluster-name"
// CertTeleportUserCertificate is the certificate of the authenticated in user.
CertTeleportUserCertificate = "x-teleport-certificate"
// extIntSuffix is the suffix common to all internal extensions.
extIntSuffix = "@teleport.internal"
// ExtIntCertType is an internal extension used to propagate cert type.
ExtIntCertType = "certtype" + extIntSuffix
// ExtIntCertTypeHost indicates a host-type certificate.
ExtIntCertTypeHost = "host" + extIntSuffix
// ExtIntCertTypeUser indicates a user-type certificate.
ExtIntCertTypeUser = "user" + extIntSuffix
// ExtIntSSHAccessPermit is an internal extension used to propagate
// the access permit for the user.
ExtIntSSHAccessPermit = "ssh-access-permit" + extIntSuffix
// ExtIntSSHJoinPermi is an internal extension used to propagate
// the join permit for the user.
ExtIntSSHJoinPermit = "ssh-join-permit" + extIntSuffix
// ExtIntProxyingPermit is an internal extension used to propagate
// the proxying permit for the user.
ExtIntProxyingPermit = "proxying-permit" + extIntSuffix
// ExtIntGitForwardingPermit is an internal extension used to propagate
// the git forwarding permit for the user.
ExtIntGitForwardingPermit = "git-forwarding-permit" + extIntSuffix
)
// IsInternalSSHExtension returns true if the extension has the internal
// extension suffix.
func IsInternalSSHExtension(extension string) bool {
return strings.HasSuffix(extension, extIntSuffix)
}
/*
* 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 utils
import (
"fmt"
"github.com/coreos/go-semver/semver"
"github.com/gravitational/trace"
)
// MeetsMinVersion returns true if gotVer is empty or at least minVer.
func MeetsMinVersion(gotVer, minVer string) bool {
if gotVer == "" {
return true // Ignore empty versions.
}
err := CheckMinVersion(gotVer, minVer)
// Non BadParameter errors are semver parsing errors.
return !trace.IsBadParameter(err)
}
// MeetsMaxVersion returns true if gotVer is empty or at most maxVer.
func MeetsMaxVersion(gotVer, maxVer string) bool {
if gotVer == "" {
return true // Ignore empty versions.
}
err := CheckMaxVersion(gotVer, maxVer)
// Non BadParameter errors are semver parsing errors.
return !trace.IsBadParameter(err)
}
// CheckMinVersion compares a version with a minimum version supported.
func CheckMinVersion(currentVersion, minVersion string) error {
currentSemver, minSemver, err := versionStringsToSemver(currentVersion, minVersion)
if err != nil {
return trace.Wrap(err)
}
if currentSemver.LessThan(*minSemver) {
return trace.BadParameter("incompatible versions: %v < %v", currentVersion, minVersion)
}
return nil
}
// CheckMaxVersion compares a version with a maximum version supported.
func CheckMaxVersion(currentVersion, maxVersion string) error {
currentSemver, maxSemver, err := versionStringsToSemver(currentVersion, maxVersion)
if err != nil {
return trace.Wrap(err)
}
if maxSemver.LessThan(*currentSemver) {
return trace.BadParameter("incompatible versions: %v > %v", currentVersion, maxVersion)
}
return nil
}
// VersionBeforeAlpha appends "-aa" to the version so that it comes before <version>-alpha.
// This ban be used to make version checks work during development.
func VersionBeforeAlpha(version string) string {
return version + "-aa"
}
// VersionWithoutPreRelease removes the prerelease suffix. Useful when showing
// teleport.MinClientSemVer(), which by default comes with the -aa prerelease.
func VersionWithoutPreRelease(version string) (string, error) {
semver, err := versionStringToSemver(version)
if err != nil {
return "", trace.Wrap(err)
}
semver.PreRelease = ""
return semver.String(), nil
}
// MinVerWithoutPreRelease compares semver strings, but skips prerelease. This allows to compare
// two versions and ignore dev,alpha,beta, etc. strings.
func MinVerWithoutPreRelease(currentVersion, minVersion string) (bool, error) {
currentSemver, minSemver, err := versionStringsToSemver(currentVersion, minVersion)
if err != nil {
return false, trace.Wrap(err)
}
// Erase pre-release string, so only version is compared.
currentSemver.PreRelease = ""
minSemver.PreRelease = ""
return !currentSemver.LessThan(*minSemver), nil
}
// MajorSemver returns the major version as a semver string.
// Ex: 13.4.3 -> 13.0.0
func MajorSemver(version string) (string, error) {
ver, err := semver.NewVersion(version)
if err != nil {
return "", trace.Wrap(err)
}
return fmt.Sprintf("%d.0.0", ver.Major), nil
}
// MajorSemverWithWildcards returns the major version as a semver string.
// Ex: 13.4.3 -> 13.x.x
func MajorSemverWithWildcards(version string) (string, error) {
ver, err := semver.NewVersion(version)
if err != nil {
return "", trace.Wrap(err)
}
return fmt.Sprintf("%d.x.x", ver.Major), nil
}
func versionStringsToSemver(ver1, ver2 string) (*semver.Version, *semver.Version, error) {
v1Semver, err := versionStringToSemver(ver1)
if err != nil {
return nil, nil, trace.Wrap(err)
}
v2Semver, err := versionStringToSemver(ver2)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return v1Semver, v2Semver, nil
}
func versionStringToSemver(ver string) (*semver.Version, error) {
semver, err := semver.NewVersion(ver)
if err != nil {
return nil, trace.Wrap(err, "unsupported version format, need semver format: %q, e.g 1.0.0", semver)
}
return semver, nil
}
/*
* 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 utils
// CaptureNBytesWriter is an io.Writer thats captures up to first n bytes
// of the incoming data in memory, and then it ignores the rest of the incoming
// data.
type CaptureNBytesWriter struct {
capture []byte
maxRemaining int
}
// NewCaptureNBytesWriter creates a new CaptureNBytesWriter.
func NewCaptureNBytesWriter(max int) *CaptureNBytesWriter {
return &CaptureNBytesWriter{
maxRemaining: max,
}
}
// Write implements io.Writer.
func (w *CaptureNBytesWriter) Write(p []byte) (int, error) {
if w.maxRemaining > 0 {
capture := p[:]
if len(capture) > w.maxRemaining {
capture = capture[:w.maxRemaining]
}
w.capture = append(w.capture, capture...)
w.maxRemaining -= len(capture)
}
// Always pretend to be successful.
return len(p), nil
}
// Bytes returns all captured bytes.
func (w CaptureNBytesWriter) Bytes() []byte {
return w.capture
}
/*
* 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 (
"bufio"
"context"
"log/slog"
"net"
"net/http"
"net/netip"
"strings"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/utils"
)
const xForwardedForHeader = "X-Forwarded-For"
// NewXForwardedForMiddleware is an HTTP middleware that overwrites client
// source address if X-Forwarded-For is set.
//
// Both hijacked conn and request context are updated. The hijacked conn can be
// used for ALPN connection upgrades or Websocket connections.
func NewXForwardedForMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
clientSrcAddr, err := parseXForwardedForHeaders(r.RemoteAddr, r.Header.Values(xForwardedForHeader))
switch {
// Skip updating client source address if no X-Forwarded-For is
// present. For example, the request may come from an internal
// network or the load balancer itself.
case trace.IsNotFound(err):
next.ServeHTTP(w, r)
// Reject the request on error.
case err != nil:
trace.WriteError(w, err)
// Serve with updated client source address.
default:
next.ServeHTTP(
responseWriterWithClientSrcAddr(r.Context(), w, clientSrcAddr),
requestWithClientSrcAddr(r, clientSrcAddr),
)
}
})
}
// parseXForwardedForHeaders returns a net.Addr from provided values of X-Forwarded-For.
//
// MDN reference:
// https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/X-Forwarded-For
//
// AWS ALB reference:
// https://docs.aws.amazon.com/elasticloadbalancing/latest/application/x-forwarded-headers.html
func parseXForwardedForHeaders(observedAddr string, xForwardedForHeaders []string) (net.Addr, error) {
switch len(xForwardedForHeaders) {
case 0:
return nil, trace.NotFound("no X-Forwarded-For headers")
case 1:
// Reject multiple IPs.
if strings.Contains(xForwardedForHeaders[0], ",") {
return nil, trace.BadParameter("expect a single IP from X-Forwarded-For but got %v", xForwardedForHeaders)
}
default:
// Reject multiple IPs.
return nil, trace.BadParameter("expect a single IP from X-Forwarded-For but got %v", xForwardedForHeaders)
}
// If forwardedAddr has a port, use that.
forwardedAddr := strings.TrimSpace(xForwardedForHeaders[0])
if ipAddrPort, err := netip.ParseAddrPort(forwardedAddr); err == nil {
return net.TCPAddrFromAddrPort(ipAddrPort), nil
}
// If forwardedAddr does not have a port, use port from observedAddr.
ipAddr, err := netip.ParseAddr(forwardedAddr)
if err != nil {
return nil, trace.BadParameter("invalid X-Forwarded-For %v: %v", xForwardedForHeaders, err)
}
var port int
if parsed, err := utils.ParseAddr(observedAddr); err == nil {
port = parsed.Port(port)
}
return net.TCPAddrFromAddrPort(netip.AddrPortFrom(ipAddr, uint16(port))), nil
}
func requestWithClientSrcAddr(r *http.Request, clientSrcAddr net.Addr) *http.Request {
ctx := authz.ContextWithClientSrcAddr(r.Context(), clientSrcAddr)
r = r.WithContext(ctx)
r.RemoteAddr = clientSrcAddr.String()
return r
}
func responseWriterWithClientSrcAddr(ctx context.Context, w http.ResponseWriter, clientSrcAddr net.Addr) http.ResponseWriter {
// Returns the original ResponseWriter if not a http.Hijacker.
_, ok := w.(http.Hijacker)
if !ok {
slog.DebugContext(ctx, "Provided ResponseWriter is not a hijacker")
return w
}
return &responseWriterWithRemoteAddr{
ResponseWriter: w,
remoteAddr: clientSrcAddr,
}
}
// responseWriterWithRemoteAddr is a wrapper of provided http.ResponseWriter
// and overwrites Hijacker interface to return a net.Conn with provided
// remoteAddr.
type responseWriterWithRemoteAddr struct {
http.ResponseWriter
remoteAddr net.Addr
}
// Hijack returns a net.Conn with provided remoteAddr.
func (r *responseWriterWithRemoteAddr) Hijack() (net.Conn, *bufio.ReadWriter, error) {
hijacker, ok := r.ResponseWriter.(http.Hijacker)
if !ok {
return nil, nil, trace.BadParameter("provided ResponseWriter is not a hijacker")
}
conn, buffer, err := hijacker.Hijack()
if err != nil {
return conn, buffer, trace.Wrap(err)
}
return utils.NewConnWithSrcAddr(conn, r.remoteAddr), buffer, nil
}
/*
* 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 implements web proxy handler that provides
// web interface to view and connect to teleport nodes
package web
import (
"bytes"
"cmp"
"context"
"crypto/tls"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"io"
"log/slog"
"math/rand/v2"
"net"
"net/http"
"net/url"
"os"
"regexp"
"slices"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
htmltemplate "github.com/DataDog/datadog-agent/pkg/template/html"
texttemplate "github.com/DataDog/datadog-agent/pkg/template/text"
gogoproto "github.com/gogo/protobuf/proto"
"github.com/google/safetext/shsprintf"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/gravitational/roundtrip"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/julienschmidt/httprouter"
"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"
"google.golang.org/grpc"
"google.golang.org/protobuf/encoding/protojson"
"google.golang.org/protobuf/types/known/timestamppb"
"github.com/gravitational/teleport"
apiclient "github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/client/webclient"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1"
notificationsv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/notifications/v1"
scopedaccessv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/access/v1"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
summarizerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/summarizer/v1"
"github.com/gravitational/teleport/api/mfa"
apitracing "github.com/gravitational/teleport/api/observability/tracing"
apissh "github.com/gravitational/teleport/api/ssh"
"github.com/gravitational/teleport/api/types"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/api/types/installers"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/api/utils/keys"
apisshutils "github.com/gravitational/teleport/api/utils/sshutils"
"github.com/gravitational/teleport/entitlements"
"github.com/gravitational/teleport/lib/auth"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/auth/moderation"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/automaticupgrades"
autoupdatelookup "github.com/gravitational/teleport/lib/autoupdate/lookup"
"github.com/gravitational/teleport/lib/client"
dbrepl "github.com/gravitational/teleport/lib/client/db/repl"
"github.com/gravitational/teleport/lib/client/sso"
"github.com/gravitational/teleport/lib/componentfeatures"
"github.com/gravitational/teleport/lib/defaults"
dtconfig "github.com/gravitational/teleport/lib/devicetrust/config"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/httplib/csrf"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/jwt"
"github.com/gravitational/teleport/lib/limiter"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/multiplexer"
"github.com/gravitational/teleport/lib/observability/tracing"
"github.com/gravitational/teleport/lib/player"
"github.com/gravitational/teleport/lib/plugin"
"github.com/gravitational/teleport/lib/proxy"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/scopes"
scopedutils "github.com/gravitational/teleport/lib/scopes/utils"
"github.com/gravitational/teleport/lib/secret"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/services/readonly"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/srv/server/installstatus"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/utils/set"
"github.com/gravitational/teleport/lib/web/app"
websession "github.com/gravitational/teleport/lib/web/session"
"github.com/gravitational/teleport/lib/web/terminal"
"github.com/gravitational/teleport/lib/web/ui"
)
const (
// SSOLoginFailureMessage is a generic error message to avoid disclosing sensitive SSO failure messages.
SSOLoginFailureMessage = "Failed to login. Please check Teleport's log for more details."
// SSOLoginFailureInvalidRedirect is a slightly specific error message for
// SSO failures related to the use of an invalid or disallowed login
// callback URL in tsh login.
SSOLoginFailureInvalidRedirect = "Failed to login due to a disallowed callback URL. Please check Teleport's log for more details."
// 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,
// to the workflow it was added, or providing unique features to those resources.
// Example values:
// - github-actions-ssh: indicates that the resource was added via the Bot GitHub Actions SSH flow
webUIFlowLabelKey = "teleport.internal/ui-flow"
// IncludedResourceModeAll describes that only requestable resources should be returned.
IncludedResourceModeRequestable = "requestable"
// IncludedResourceModeAll describes that all resources, requestable and available, should be returned.
IncludedResourceModeAll = "all"
// DefaultFeatureWatchInterval is the default time in which the feature watcher
// should ping the auth server to check for updated features
DefaultFeatureWatchInterval = time.Minute * 5
// findEndpointCacheTTL is the cache TTL for the find endpoint generic answer.
// This cache is here to protect against accidental or intentional DDoS, the TTL must be low to quickly reflect
// cluster configuration changes.
findEndpointCacheTTL = 10 * time.Second
// DefaultAgentUpdateJitterSeconds is the default jitter agents should wait before updating.
DefaultAgentUpdateJitterSeconds = 60
)
// healthCheckAppServerFunc defines a function used to perform a health check
// to AppServer that can handle application requests (based on cluster name and
// public address).
type healthCheckAppServerFunc func(ctx context.Context, appName, publicAddr, clusterName string) error
// Handler is HTTP web proxy handler
type Handler struct {
logger *slog.Logger
sync.Mutex
httprouter.Router
cfg Config
auth *sessionCache
clock clockwork.Clock
limiter *limiter.RateLimiter
highLimiter *limiter.RateLimiter
healthCheckAppServer healthCheckAppServerFunc
// sshPort specifies the SSH proxy port extracted
// from configuration
sshPort string
// userConns tracks amount of current active connections with user certificates.
userConns atomic.Int32
// clusterFeatures contain flags for supported and unsupported features.
clusterFeatures proto.Features
// nodeWatcher is a services.NodeWatcher used by Assist to lookup nodes from
// the proxy's cache and get nodes in real time.
nodeWatcher *services.GenericWatcher[types.Server, readonly.Server]
// appServerWatcher ia a app server watcher to speed up app look up.
appServerWatcher *services.GenericWatcher[types.AppServer, readonly.AppServer]
// tracer is used to create spans.
tracer oteltrace.Tracer
// findEndpointCache is used to cache the find endpoint answer. As this endpoint is unprotected and has high
// rate-limits, each call must cause minimal work. The cached answer can be modulated after, for example if the
// caller specified its Automatic Updates UUID or group.
findEndpointCache *utils.FnCache
autoUpdateResolver *autoupdatelookup.Resolver
accessGraphHandler http.Handler
// webSessionRootClientDialOptions contains additional gRPC dial options for
// root clients created for web sessions. Used for testing.
webSessionRootClientDialOptions []grpc.DialOption
}
// HandlerOption is a functional argument - an option that can be passed
// to NewHandler function
type HandlerOption func(h *Handler) error
// SetClock sets the clock on a handler
func SetClock(clock clockwork.Clock) HandlerOption {
return func(h *Handler) error {
h.clock = clock
return nil
}
}
// WithWebSessionRootClientDialOption adds a [grpc.DialOption] to root clients
// created for web sessions. It is intended for tests.
func WithWebSessionRootClientDialOption(opt grpc.DialOption) HandlerOption {
return func(h *Handler) error {
h.webSessionRootClientDialOptions = append(h.webSessionRootClientDialOptions, opt)
return nil
}
}
type ProxySettingsGetter interface {
GetProxySettings(ctx context.Context) (*webclient.ProxySettings, error)
}
// PresenceChecker is a function that executes an MFA prompt to enforce
// that a user is present.
type PresenceChecker = func(ctx context.Context, term io.Writer, maintainer client.PresenceMaintainer, sessionID string, mfaCeremony *mfa.Ceremony, opts ...client.PresenceOption) error
// Config represents web handler configuration parameters
type Config struct {
// PluginRegistry handles plugin registration
PluginRegistry plugin.Registry
// Proxy provides a means to look up clusters.
Proxy reversetunnelclient.ClusterGetter
// AuthServers is a list of auth servers this proxy talks to
AuthServers utils.NetAddr
// ProxyClient is a client that authenticated as proxy
ProxyClient authclient.ClientI
// ProxySSHAddr points to the SSH address of the proxy
ProxySSHAddr utils.NetAddr
// ProxyKubeAddr points to the Kube address of the proxy
ProxyKubeAddr utils.NetAddr
// ProxyWebAddr points to the web (HTTPS) address of the proxy
ProxyWebAddr utils.NetAddr
// ProxyPublicAddr contains web proxy public addresses.
ProxyPublicAddrs []utils.NetAddr
// ProxyGroupID is reverse tunnel group ID, used by reverse tunnel agents
// in proxy peering mode.
ProxyGroupID string
// GetProxyClientCertificate returns the proxy client certificate.
GetProxyClientCertificate func() (*tls.Certificate, error)
// CipherSuites is the list of cipher suites Teleport suppports.
CipherSuites []uint16
// FIPS mode means Teleport started in a FedRAMP/FIPS compliant
// configuration.
FIPS bool
// InsecureMode defines whether insecure connections are allowed.
InsecureMode bool
// Modules define the build type, entitlements and licensed features.
Modules modules.Modules
// AccessPoint holds a cache to the Auth Server.
AccessPoint authclient.ProxyAccessPoint
// Emitter is event emitter
Emitter apievents.Emitter
// HostUUID is the UUID of this process.
HostUUID string
// Context is used to signal process exit.
Context context.Context
// StaticFS optionally specifies the HTTP file system to use.
// Enables web UI if set.
StaticFS http.FileSystem
// CachedSessionLingeringThreshold specifies the time the session will linger
// in the cache before getting purged after it has expired.
// Defaults to cachedSessionLingeringThreshold if unspecified.
CachedSessionLingeringThreshold *time.Duration
// ClusterFeatures contains flags for supported/unsupported features.
ClusterFeatures proto.Features
// ScopesFeatures dictates which scoped components are enabled for this server.
ScopesFeatures scopes.Features
// ProxySettings allows fetching the current proxy settings.
ProxySettings ProxySettingsGetter
// MinimalReverseTunnelRoutesOnly mode handles only the endpoints required for
// a reverse tunnel agent to establish a connection.
MinimalReverseTunnelRoutesOnly bool
// PublicProxyAddr is used to template the public proxy address
// into the installer script responses
PublicProxyAddr string
// ALPNHandler is the ALPN connection handler for handling upgraded ALPN
// connection through an HTTP upgrade call.
//
// It’s also used in scenarios where the Proxy needs to dial to itself (e.g.
// database access via ws), but the handler can directly forward the traffic
// to the ALPN router without initiating a new connection.
ALPNHandler ConnectionHandler
// TraceClient is used to forward spans to the upstream collector for the UI
TraceClient otlptrace.Client
// Router is used to route ssh sessions to hosts
Router *proxy.Router
// SessionControl is used to determine if users are
// allowed to spawn new sessions
SessionControl SessionController
// PROXYSigner is used to sign PROXY header and securely propagate client IP information
PROXYSigner multiplexer.PROXYHeaderSigner
// TracerProvider generates tracers to create spans with
TracerProvider oteltrace.TracerProvider
// HealthCheckAppServer is a function that checks if the proxy can handle
// application requests.
HealthCheckAppServer healthCheckAppServerFunc
// UI provides config options for the web UI
UI webclient.UIConfig
// NodeWatcher is a services.NodeWatcher used by Assist to lookup nodes from
// the proxy's cache and get nodes in real time.
NodeWatcher *services.GenericWatcher[types.Server, readonly.Server]
// AppServerWatcher ia a app server watcher to speed up app look up.
AppServerWatcher *services.GenericWatcher[types.AppServer, readonly.AppServer]
// PresenceChecker periodically runs the mfa ceremony for moderated
// sessions.
PresenceChecker PresenceChecker
// AccessGraphAddr is the address of the Access Graph service GRPC API
AccessGraphAddr utils.NetAddr
// AutomaticUpgradesChannels is a map of all version channels used by the
// proxy built-in version server to retrieve target versions. This is part
// of the automatic upgrades.
AutomaticUpgradesChannels automaticupgrades.Channels
// IntegrationAppHandler handles App Access requests which use an Integration.
IntegrationAppHandler app.ServerHandler
// FeatureWatchInterval is the interval between pings to the auth server
// to fetch new cluster features
FeatureWatchInterval time.Duration
// DatabaseREPLRegistry is used for retrieving database REPL.
DatabaseREPLRegistry dbrepl.REPLRegistry
}
// SetDefaults ensures proper default values are set if
// not provided.
func (c *Config) SetDefaults() {
c.ProxyClient = auth.WithGithubConnectorConversions(c.ProxyClient)
if c.TracerProvider == nil {
c.TracerProvider = tracing.NoopProvider()
}
if c.PresenceChecker == nil {
c.PresenceChecker = client.RunDefaultPresenceTask
}
if c.AutomaticUpgradesChannels == nil {
c.AutomaticUpgradesChannels = automaticupgrades.Channels{}
}
// TODO(tross): remove this when modules are injected properly.
if c.Modules == nil {
c.Modules = modules.GetModules()
}
c.ProxyGroupID = cmp.Or(c.ProxyGroupID, os.Getenv("TELEPORT_UNSTABLE_PROXYGROUP_ID"))
c.FeatureWatchInterval = cmp.Or(c.FeatureWatchInterval, DefaultFeatureWatchInterval)
}
type APIHandler struct {
handler *Handler
// appHandler is a http.Handler to forward requests to applications.
appHandler *app.Handler
}
// ConnectionHandler defines a function for serving incoming connections.
type ConnectionHandler func(ctx context.Context, conn net.Conn) error
func (h *APIHandler) handlePreflight(w http.ResponseWriter, r *http.Request) {
raddr, err := utils.ParseAddr(r.Host)
if err != nil {
return
}
publicAddr := raddr.Host()
servers, err := h.handler.appServerWatcher.CurrentResourcesWithFilter(r.Context(), app.MatchPublicAddr(publicAddr))
if err != nil {
h.handler.logger.InfoContext(r.Context(), "failed to match application with public addr", "public_addr", publicAddr)
return
}
if len(servers) == 0 {
h.handler.logger.InfoContext(r.Context(), "failed to match application with public addr", "public_addr", publicAddr)
return
}
foundApp := servers[rand.N(len(servers))].GetApp()
corsPolicy := foundApp.GetCORS()
if corsPolicy == nil {
return
}
origin := r.Header.Get("Origin")
// The Access-Control-Allow-Origin can only include one origin or a wildcard. However,
// any request which includes credentials _must_ return an origin and not a wildcard.
// https://developer.mozilla.org/en-US/docs/Web/HTTP/CORS#sect2
if slices.Contains(corsPolicy.AllowedOrigins, "*") || slices.Contains(corsPolicy.AllowedOrigins, origin) {
w.Header().Set("Access-Control-Allow-Origin", origin)
} else {
return
}
if len(corsPolicy.AllowedMethods) > 0 {
w.Header().Set("Access-Control-Allow-Methods", strings.Join(corsPolicy.AllowedMethods, ","))
}
// This is a list of headers that are allowed in the spec. Wildcards are allowed.
// Note: "Authorization" headers must be explicitly listed and cannot be wildcarded
// https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/Access-Control-Allow-Headers#sect2
if len(corsPolicy.AllowedHeaders) > 0 {
w.Header().Set("Access-Control-Allow-Headers", strings.Join(corsPolicy.AllowedHeaders, ","))
}
if len(corsPolicy.ExposedHeaders) > 0 {
w.Header().Set("Access-Control-Expose-Headers", strings.Join(corsPolicy.ExposedHeaders, ","))
}
// The only valid value for this header is "true", so we will only set it if configured to true
if corsPolicy.AllowCredentials {
w.Header().Set("Access-Control-Allow-Credentials", "true")
}
// This will allow preflight responses to be cached for the specified duration
if corsPolicy.MaxAge > 0 {
w.Header().Set("Access-Control-Max-Age", fmt.Sprintf("%d", corsPolicy.MaxAge))
}
w.WriteHeader(http.StatusOK)
}
// Check if this request should be forwarded to an application handler to
// be handled by the UI and handle the request appropriately.
func (h *APIHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// If the request is for the Access Graph API, forward to the access graph handler.
// This handler is only setup for enterprise but for OSS and non licensed clusters,
// it will return an error indicating that the feature is not enabled.
if isAccessGraphAPIRequest(r) {
if h.handler.accessGraphHandler == nil {
trace.WriteError(w, trace.NotImplemented("access graph is not enabled"))
return
}
h.handler.accessGraphHandler.ServeHTTP(w, r)
return
}
// If the request is either to the fragment authentication endpoint, a DBSC
// endpoint, or if the request has a session cookie or a client cert, forward
// to application handlers. If the request is requesting a FQDN that is not of
// the proxy, redirect to application launcher.
if h.appHandler != nil && shouldForwardToAppHandler(r) {
h.appHandler.ServeHTTP(w, r)
return
}
// Build the proxy address list for app routing. When
// proxy_service.public_addr is not configured, only activate the
// cluster-name fallback if the request host is a subdomain of the
// cluster name. This avoids misclassifying requests that arrive on
// a different hostname (e.g. behind a load balancer) as app requests.
proxyAddrs := h.handler.cfg.ProxyPublicAddrs
if len(proxyAddrs) == 0 && h.appHandler != nil {
clusterName := h.handler.auth.clusterName
raddr, err := utils.ParseAddr(r.Host)
if err == nil && clusterName != "" && strings.HasSuffix(raddr.Host(), "."+clusterName) {
port := raddr.Port(443)
host := net.JoinHostPort(clusterName, strconv.Itoa(port))
proxyAddrs = []utils.NetAddr{{Addr: host}}
}
}
// if the request is for an app, passthrough OPTIONS requests to the app handler
redir, ok := app.HasName(r, proxyAddrs)
if ok && r.Method == http.MethodOptions {
h.handlePreflight(w, r)
return
}
// Only try to redirect if the handler is serving the full Web API.
if !h.handler.cfg.MinimalReverseTunnelRoutesOnly && ok {
http.Redirect(w, r, redir, http.StatusFound)
return
}
// Serve the Web UI.
h.handler.ServeHTTP(w, r)
}
func shouldForwardToAppHandler(r *http.Request) bool {
return app.HasFragment(r) ||
app.IsDBSCRequest(r) ||
app.HasSessionCookie(r) ||
app.HasClientCert(r) ||
app.IsHTTPSTunnelConn(r)
}
// SetAccessGraphHandler sets the handler used to serve Access Graph API
// requests authenticated via a client TLS certificate (RouteToApp.AccessGraph=true).
// Only called by the enterprise plugin; remains nil on OSS clusters.
func (h *Handler) SetAccessGraphHandler(handler http.Handler) {
h.accessGraphHandler = handler
}
// isAccessGraphAPIRequest returns true if the request is authenticated with a client TLS certificate
// with UsageAccessGraphAPIOnly set in the certificate's usage, indicating that the request is
// intended for the Access Graph API.
func isAccessGraphAPIRequest(r *http.Request) bool {
if r.TLS == nil || len(r.TLS.PeerCertificates) == 0 {
return false
}
cert := r.TLS.PeerCertificates[0]
identity, err := tlsca.FromSubject(cert.Subject, cert.NotAfter)
if err != nil {
return false
}
return slices.Contains(identity.Usage, teleport.UsageAccessGraphAPIOnly)
}
// HandleConnection handles connections from plain TCP applications.
func (h *APIHandler) HandleConnection(ctx context.Context, conn net.Conn) error {
return h.appHandler.HandleConnection(ctx, conn)
}
func (h *APIHandler) Close() error {
return h.handler.Close()
}
// NewHandler returns a new instance of web proxy handler
func NewHandler(cfg Config, opts ...HandlerOption) (*APIHandler, error) {
cfg.SetDefaults()
h := &Handler{
cfg: cfg,
logger: slog.Default().With(teleport.ComponentKey, teleport.ComponentWeb),
clock: clockwork.NewRealClock(),
clusterFeatures: cfg.ClusterFeatures,
healthCheckAppServer: cfg.HealthCheckAppServer,
tracer: cfg.TracerProvider.Tracer(teleport.ComponentWeb),
}
if automaticUpgrades(cfg.ClusterFeatures) && h.cfg.AutomaticUpgradesChannels == nil {
h.cfg.AutomaticUpgradesChannels = automaticupgrades.Channels{}
}
// for properly handling url-encoded parameter values.
h.UseRawPath = true
for _, o := range opts {
if err := o(h); err != nil {
return nil, trace.Wrap(err)
}
}
// We create the cache after applying the options to make sure we use the fake clock if it was passed.
findCache, err := utils.NewFnCache(utils.FnCacheConfig{
TTL: findEndpointCacheTTL,
Clock: h.clock,
Context: cfg.Context,
ReloadOnErr: false,
})
if err != nil {
return nil, trace.Wrap(err, "creating /find cache")
}
h.findEndpointCache = findCache
autoUpdateResolver, err := autoupdatelookup.NewResolver(
autoupdatelookup.Config{
RolloutGetter: cfg.AccessPoint,
CMCGetter: cfg.ProxyClient,
Channels: h.cfg.AutomaticUpgradesChannels,
Log: h.logger,
Clock: h.clock,
Context: h.cfg.Context,
})
if err != nil {
return nil, trace.Wrap(err, "creating autoupdate resolver")
}
h.autoUpdateResolver = autoUpdateResolver
// We create the cache after applying the options to make sure we use the fake clock if it was passed.
sessionLingeringThreshold := cachedSessionLingeringThreshold
if cfg.CachedSessionLingeringThreshold != nil {
sessionLingeringThreshold = *cfg.CachedSessionLingeringThreshold
}
sessionCache, err := newSessionCache(h.cfg.Context, sessionCacheOptions{
proxyClient: cfg.ProxyClient,
accessPoint: cfg.AccessPoint,
scopedRoleReader: cfg.AccessPoint.ScopedRoleReader(),
servers: []utils.NetAddr{cfg.AuthServers},
cipherSuites: cfg.CipherSuites,
clock: h.clock,
sessionLingeringThreshold: sessionLingeringThreshold,
proxySigner: cfg.PROXYSigner,
logger: h.logger,
buildType: cfg.Modules.BuildType(),
rootClientDialOptions: h.webSessionRootClientDialOptions,
})
if err != nil {
return nil, trace.Wrap(err)
}
h.auth = sessionCache
sshPortValue := strconv.Itoa(defaults.SSHProxyListenPort)
if cfg.ProxySSHAddr.String() != "" {
_, sshPort, err := net.SplitHostPort(cfg.ProxySSHAddr.String())
if err != nil {
h.logger.WarnContext(h.cfg.Context, "Invalid SSH proxy address, will use default port",
"error", err,
"ssh_proxy_addr", logutils.StringerAttr(&cfg.ProxySSHAddr),
"default_port", defaults.SSHProxyListenPort,
)
} else {
sshPortValue = sshPort
}
}
h.sshPort = sshPortValue
// rateLimiter is used to limit unauthenticated challenge generation for
// passwordless and for unauthenticated metrics.
h.limiter, err = limiter.NewRateLimiter(limiter.Config{
Rates: []limiter.Rate{
{
Period: defaults.LimiterPeriod,
Average: defaults.LimiterAverage,
Burst: defaults.LimiterBurst,
},
},
MaxConnections: defaults.LimiterMaxConnections,
})
if err != nil {
return nil, trace.Wrap(err)
}
// highLimiter is used for endpoints which are only CPU constrained and require high request rates
h.highLimiter, err = limiter.NewRateLimiter(limiter.Config{
Rates: []limiter.Rate{
{
Period: defaults.LimiterHighPeriod,
Average: defaults.LimiterHighAverage,
Burst: defaults.LimiterHighBurst,
},
},
MaxConnections: defaults.LimiterMaxConnections,
})
if err != nil {
return nil, trace.Wrap(err)
}
if cfg.MinimalReverseTunnelRoutesOnly {
h.bindMinimalEndpoints()
} else {
h.bindDefaultEndpoints()
}
// serve the web UI from the embedded filesystem
var indexPage *htmltemplate.Template
// we will set our etag based on the teleport version and
// the webasset app hash if available. The version only will not
// suffice as it can cause incorrect caching for local development.
// The hash of the webasset app.js is used to ensure that builds at
// different times or different OSes will be the same and not cause
// cache invalidation for production users. For example, using a timestamp
// at build time would cause different OS builds to be different, and timestamps
// at process start would mean multiple proxies would serving different etags)
etag := fmt.Sprintf("W/%q", teleport.Version)
if cfg.StaticFS != nil {
index, err := cfg.StaticFS.Open("/index.html")
if err != nil {
h.logger.ErrorContext(h.cfg.Context, "Failed to open index file", "error", err)
return nil, trace.Wrap(err)
}
defer index.Close()
indexContent, err := io.ReadAll(index)
if err != nil {
return nil, trace.ConvertSystemError(err)
}
indexPage, err = htmltemplate.New("index").Parse(string(indexContent))
if err != nil {
return nil, trace.BadParameter("failed parsing index.html template: %v", err)
}
h.Handle("GET", "/robots.txt", httplib.MakeHandler(serveRobotsTxt))
etagFromAppHash, err := readEtagFromAppHash(cfg.StaticFS)
if err != nil {
h.logger.ErrorContext(h.cfg.Context, "Could not read apphash from embedded webassets. Using version only as ETag for Web UI assets", "error", err)
} else {
etag = etagFromAppHash
}
}
// This endpoint is used both by Web UI and Connect.
h.Handle("GET", "/web/config.js", h.WithUnauthenticatedLimiter(h.getWebConfig))
if cfg.NodeWatcher != nil {
h.nodeWatcher = cfg.NodeWatcher
}
if cfg.AppServerWatcher != nil {
h.appServerWatcher = cfg.AppServerWatcher
}
const v1Prefix = "/v1"
notFoundRoutingHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// Request is going to the API?
// If no routes were matched, it could be because it's a path with `v1` prefix
// (eg: the Teleport web app will call "most" endpoints with v1 prefixed).
//
// `v1` paths are not defined with `v1` prefix. If the path turns out to be prefixed
// with `v1`, it will be stripped and served again. Historically, that's how it started
// and should be kept that way to prevent breakage.
//
// v2+ prefixes will be expected by both caller and definition and will not be stripped.
if strings.HasPrefix(r.URL.Path, v1Prefix) {
pathParts := strings.Split(r.URL.Path, "/")
if len(pathParts) > 2 {
// check against known second part of path to ensure we
// aren't allowing paths like /v1/v2/webapi
// part[0] is empty space from leading slash "/"
// part[1] is the prefix "v1"
switch pathParts[2] {
case "webapi", "enterprise", "scripts", ".well-known", "workload-identity", "web":
http.StripPrefix(v1Prefix, h).ServeHTTP(w, r)
return
}
}
httplib.RouteNotFoundResponse(r.Context(), w)
return
}
// request is going to the web UI
if cfg.StaticFS == nil {
httplib.RouteNotFoundResponse(r.Context(), w)
return
}
// redirect to "/web" when someone hits "/"
if r.URL.Path == "/" {
app.SetRedirectPageHeaders(w.Header(), "")
http.Redirect(w, r, "/web", http.StatusFound)
return
}
// serve Web UI:
if strings.HasPrefix(r.URL.Path, "/web/app") {
// Check if the incoming request wants to check the version
// and if the version has not changed, return a Not Modified response
if match := r.Header.Get("If-None-Match"); match == etag {
w.WriteHeader(http.StatusNotModified)
return
}
fs := http.FileServer(cfg.StaticFS)
fs = makeBrotliHandler(fs, cfg.StaticFS)
fs = makeCacheHandler(fs, etag)
http.StripPrefix("/web", fs).ServeHTTP(w, r)
} else if strings.HasPrefix(r.URL.Path, "/web/") || r.URL.Path == "/web" {
csrfToken, err := csrf.AddCSRFProtection(w, r)
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to generate CSRF token", "error", err)
}
// Ignore errors here, as unauthenticated requests for index.html are common - the user might
// not have logged in yet, or their session may have expired.
// The web app will show them the login page in this case.
session, _ := h.authenticateWebSession(w, r)
session.XCSRF = csrfToken
httplib.SetNoCacheHeaders(w.Header())
features := h.GetClusterFeatures()
httplib.SetIndexContentSecurityPolicy(w.Header(), features.GetIsStripeManaged(), r.URL.Path)
if err := indexPage.Execute(w, session); err != nil {
h.logger.ErrorContext(r.Context(), "Failed to execute index page template", "error", err)
}
} else {
httplib.RouteNotFoundResponse(r.Context(), w)
return
}
})
h.NotFound = notFoundRoutingHandler
if cfg.PluginRegistry != nil {
if err := cfg.PluginRegistry.RegisterProxyWebHandlers(h); err != nil {
return nil, trace.Wrap(err)
}
}
// Create application specific handler. This handler handles sessions and
// forwarding for application access.
var appHandler *app.Handler
if !cfg.MinimalReverseTunnelRoutesOnly {
appHandler, err = app.NewHandler(cfg.Context, &app.HandlerConfig{
Clock: h.clock,
AuthClient: cfg.ProxyClient,
AccessPoint: cfg.AccessPoint,
ClusterGetter: cfg.Proxy,
CipherSuites: cfg.CipherSuites,
ProxyPublicAddrs: cfg.ProxyPublicAddrs,
IntegrationAppHandler: cfg.IntegrationAppHandler,
})
if err != nil {
return nil, trace.Wrap(err)
}
if h.healthCheckAppServer == nil {
h.healthCheckAppServer = appHandler.HealthCheckAppServer
}
}
go h.startFeatureWatcher(h.cfg.Context)
return &APIHandler{
handler: h,
appHandler: appHandler,
}, nil
}
type webSession struct {
Session string
XCSRF string
}
func (h *Handler) authenticateWebSession(w http.ResponseWriter, r *http.Request) (webSession, error) {
ctx, err := h.AuthenticateRequest(w, r, false /* validate bearer token */)
if err != nil {
return webSession{}, trace.Wrap(err)
}
resp, err := newSessionResponse(r.Context(), ctx)
if err != nil {
return webSession{}, trace.Wrap(err)
}
out, err := json.Marshal(resp)
if err != nil {
return webSession{}, trace.Wrap(err)
}
return webSession{
Session: base64.StdEncoding.EncodeToString(out),
}, nil
}
// bindMinimalEndpoints binds only the endpoints required for a reverse tunnel
// agent to establish a connection.
func (h *Handler) bindMinimalEndpoints() {
// find is like ping, but is faster because it is optimized for servers
// and does not fetch the data that servers don't need, e.g.
// OIDC connectors and auth preferences
// Note that find is a unique endpoint that requires high request rates
// sometimes through NATs and thus should not be rate limited by IP.
h.GET("/webapi/find", httplib.MakeHandler(h.find))
// Issue host credentials.
h.POST("/webapi/host/credentials", h.WithUnauthenticatedHighLimiter(h.hostCredentials))
}
// bindDefaultEndpoints binds the default endpoints for the web API.
func (h *Handler) bindDefaultEndpoints() {
h.bindMinimalEndpoints()
// ping endpoint is used to check if the server is up. the /webapi/ping
// endpoint returns the default authentication method and configuration that
// the server supports. the /webapi/ping/:connector endpoint can be used to
// query the authentication configuration for a specific connector.
h.GET("/webapi/ping", httplib.MakeHandler(h.ping))
h.GET("/webapi/ping/:connector", h.WithUnauthenticatedHighLimiter(h.pingWithConnector))
// Unauthenticated access to JWT public keys.
h.GET("/.well-known/jwks.json", h.WithUnauthenticatedHighLimiter(h.wellKnownJWKS))
// Unauthenticated access to the message of the day
h.GET("/webapi/motd", h.WithHighLimiter(h.motd))
// Unauthenticated access to retrieving the script used to install Teleport
h.GET("/webapi/scripts/installer/:name", h.WithLimiter(h.installer))
// Forwards traces to the configured upstream collector
h.POST("/webapi/traces", h.WithAuth(h.traces))
// App sessions
h.POST("/webapi/sessions/app", h.WithAuth(h.createAppSession))
// Web sessions
h.POST("/webapi/sessions/web", h.WithLimiter(h.createWebSession))
h.DELETE("/webapi/sessions/web", h.WithAuth(h.deleteWebSession))
h.POST("/webapi/sessions/web/renew", h.WithAuth(h.renewWebSession))
h.POST("/webapi/users", h.WithAuth(h.createUserHandle))
h.PUT("/webapi/users", h.WithAuth(h.updateUserHandle))
// TODO(rudream): DELETE IN V21.0.0
// MUST delete with related code found in web/packages/teleport/src/services/user/user.ts(fetchUsers)
h.GET("/webapi/users", h.WithAuth(h.getUsersHandle))
// The v2 version of this endpoint is paginated.
h.GET("/v2/webapi/users", h.WithAuth(h.listUsersHandle))
h.DELETE("/webapi/users/:username", h.WithAuth(h.deleteUserHandle))
// We have an overlap route here, please see godoc of handleGetUserOrResetToken
// h.GET("/webapi/users/:username", h.WithAuth(h.getUserHandle))
// h.GET("/webapi/users/password/token/:token", h.WithLimiter(h.getResetPasswordTokenHandle))
h.GET("/webapi/users/*wildcard", h.handleGetUserOrResetToken)
h.PUT("/webapi/users/password/token", h.WithLimiter(h.changeUserAuthentication))
h.PUT("/webapi/users/password", h.WithAuth(h.changePassword))
h.POST("/webapi/users/password/token", h.WithAuth(h.createResetPasswordToken))
h.POST("/webapi/users/privilege/token", h.WithAuth(h.createPrivilegeTokenHandle))
h.POST("/webapi/headless/login", h.WithUnauthenticatedLimiter(h.headlessLogin))
// list available sites
h.GET("/webapi/sites", h.WithAuth(h.getClusters))
// Site specific API
// get site info
h.GET("/webapi/sites/:site/info", h.WithClusterAuth(h.getClusterInfo))
// get namespaces
h.GET("/webapi/sites/:site/namespaces", h.WithClusterAuth(h.getSiteNamespaces))
// get unified resources
h.GET("/webapi/sites/:site/resources", h.WithClusterAuth(h.clusterUnifiedResourcesGet))
// get nodes
h.GET("/webapi/sites/:site/nodes", h.WithClusterAuth(h.clusterNodesGet))
h.POST("/webapi/sites/:site/nodes", h.WithClusterAuth(h.handleNodeCreate))
h.GET("/webapi/sites/:site/instances", h.WithClusterAuth(h.clusterUnifiedInstancesGet))
// get login alerts
h.GET("/webapi/sites/:site/alerts", h.WithClusterAuth(h.clusterLoginAlertsGet))
// lock interactions
// TODO(nicholasmarais1158): DELETE IN 20.0.0 - Replaced by /v2/webapi/sites/:site/locks endpoint
h.GET("/webapi/sites/:site/locks", h.WithClusterAuth(h.getClusterLocks))
h.GET("/v2/webapi/sites/:site/locks", h.WithClusterAuth(h.getClusterLocksV2))
h.PUT("/webapi/sites/:site/locks", h.WithClusterAuth(h.createClusterLock))
h.DELETE("/webapi/sites/:site/locks/:uuid", h.WithClusterAuth(h.deleteClusterLock))
// active sessions handlers
h.GET("/webapi/sites/:site/connect/ws", h.WithClusterAuthWebSocket(h.siteNodeConnect)) // connect to an active session (via websocket, with auth over websocket)
h.GET("/webapi/sites/:site/sessions", h.WithClusterAuth(h.clusterActiveAndPendingSessionsGet)) // get list of active and pending sessions
h.GET("/webapi/sites/:site/kube/exec/ws", h.WithClusterAuthWebSocket(h.podConnect)) // connect to a pod with exec (via websocket, with auth over websocket)
h.GET("/webapi/sites/:site/db/exec/ws", h.WithClusterAuthWebSocket(h.dbConnect))
// Audit events handlers.
// TODO (avatus): delete in v21
// Deprecated: Use the v2 endpoint instead.
//
// clusterSearchEvents handles audit event retrieval for a given site.
// This legacy endpoint returns event listings without advanced search capabilities.
// Prefer using /v2/webapi/sites/:site/events/search for full query-based filtering.
h.GET("/webapi/sites/:site/events/search", h.WithClusterAuth(h.clusterSearchEvents)) // search site events
// clusterSearchEventsV2 handles audit event retrieval for a given site with support for
// advanced search filters and query parameters.
h.GET("/v2/webapi/sites/:site/events/search", h.WithClusterAuth(h.clusterSearchEventsV2)) // search site events
h.GET("/webapi/sites/:site/events/search/sessions", h.WithClusterAuth(h.clusterSearchSessionEvents)) // search site session events
h.GET("/webapi/sites/:site/ttyplayback/:sid", h.WithClusterAuth(h.ttyPlaybackHandle))
h.GET("/webapi/sites/:site/sessionlength/:sid", h.WithClusterAuth(h.sessionLengthHandle))
// scp file transfer
h.GET("/webapi/sites/:site/nodes/:server/:login/scp", h.WithClusterAuth(h.transferFile))
h.POST("/webapi/sites/:site/nodes/:server/:login/scp", h.WithClusterAuth(h.transferFile))
// Sign required files to set up mTLS using the db format.
h.POST("/webapi/sites/:site/sign/db", h.WithProvisionTokenAuth(h.signDatabaseCertificate))
// Returns the CA Certs
// Deprecated, use the `webapi/auth/export` endpoint.
// Returning other clusters (trusted cluster) CA certs would leak whether the TrustedCluster exists or not.
// Given that this is a public/unauthorized endpoint, we should refrain from exposing that kind of information.
h.GET("/webapi/sites/:site/auth/export", h.authExportPublic)
h.GET("/webapi/auth/export", h.authExportPublic)
// join token handlers
h.PUT("/webapi/tokens/yaml", h.WithAuth(h.updateTokenYAML))
// used for creating a new token
h.POST("/webapi/tokens", h.WithAuth(h.upsertTokenHandle))
// used for updating a token
h.PUT("/webapi/tokens", h.WithAuth(h.upsertTokenHandle))
// TODO(kimlisa): DELETE IN 19.0 - Replaced by /v2/webapi/token endpoint
// MUST delete with related code found in web/packages/teleport/src/services/joinToken/joinToken.ts(fetchJoinToken)
h.POST("/webapi/token", h.WithAuth(h.createTokenForDiscoveryHandle))
// used for creating tokens used during guided discover flows
// v2 endpoint processes "suggestedLabels" field
h.POST("/v2/webapi/token", h.WithAuth(h.createTokenForDiscoveryHandle))
h.GET("/webapi/tokens", h.WithAuth(h.getTokens))
// used to retrieve a paginated list of tokens. Items can be filtered by roles and bot name.
h.GET("/v2/webapi/tokens", h.WithAuth(h.listProvisionTokens))
h.DELETE("/webapi/tokens", h.WithAuth(h.deleteToken))
// install script, the ':token' wildcard is a hack to make the router happy and support
// the token-less route "/scripts/install.sh".
// h.installScriptHandle Will reject any unknown sub-route.
h.GET("/scripts/:token", h.WithHighLimiter(h.installScriptHandle))
// join scripts
h.GET("/scripts/:token/install-node.sh", h.WithLimiter(h.getNodeJoinScriptHandle))
h.GET("/scripts/:token/install-app.sh", h.WithLimiter(h.getAppJoinScriptHandle))
h.GET("/scripts/:token/install-database.sh", h.WithLimiter(h.getDatabaseJoinScriptHandle))
// Discovery installation script requires a query param to define the DiscoveryGroup:
// ?discoveryGroup=<group name>
h.GET("/scripts/:token/install-discovery.sh", h.WithLimiter(h.getDiscoveryJoinScriptHandle))
// web context
h.GET("/webapi/sites/:site/context", h.WithClusterAuth(h.getUserContext))
// Database access handlers.
h.GET("/webapi/sites/:site/databases", h.WithClusterAuth(h.clusterDatabasesGet))
h.POST("/webapi/sites/:site/databases", h.WithClusterAuth(h.handleDatabaseCreateOrOverwrite))
h.PUT("/webapi/sites/:site/databases/:database", h.WithClusterAuth(h.handleDatabasePartialUpdate))
h.GET("/webapi/sites/:site/databases/:database", h.WithClusterAuth(h.clusterDatabaseGet))
h.GET("/webapi/sites/:site/databases/:database/iam/policy", h.WithClusterAuth(h.handleDatabaseGetIAMPolicy))
h.GET("/webapi/scripts/databases/configure/sqlserver/:token/configure-ad.ps1", httplib.MakeHandler(h.sqlServerConfigureADScriptHandle))
// DatabaseService handlers
h.GET("/webapi/sites/:site/databaseservices", h.WithClusterAuth(h.clusterDatabaseServicesList))
// Database server handlers
h.GET("/webapi/sites/:site/databaseservers", h.WithClusterAuth(h.clusterDatabaseServersList))
// Kube access handlers.
h.GET("/webapi/sites/:site/kubernetes", h.WithClusterAuth(h.clusterKubesGet))
h.GET("/webapi/sites/:site/kubernetes/resources", h.WithClusterAuth(h.clusterKubeResourcesGet))
h.GET("/webapi/sites/:site/kubernetesservers", h.WithClusterAuth(h.clusterKubeServersList))
// Github connector handlers
h.GET("/webapi/github/login/web", h.WithRedirect(h.githubLoginWeb))
h.GET("/webapi/github/callback", h.WithMetaRedirect(h.githubCallback))
h.POST("/webapi/github/login/console", h.WithLimiter(h.githubLoginConsole))
// MFA public endpoints.
h.POST("/webapi/sites/:site/mfa/required", h.WithClusterAuth(h.isMFARequired))
h.POST("/webapi/mfa/login/begin", h.WithLimiter(h.mfaLoginBegin))
h.POST("/webapi/mfa/login/finish", h.WithLimiter(h.mfaLoginFinish))
h.POST("/webapi/mfa/login/finishsession", h.WithLimiter(h.mfaLoginFinishSession))
h.DELETE("/webapi/mfa/token/:token/devices/:devicename", h.WithLimiter(h.deleteMFADeviceWithTokenHandle))
h.GET("/webapi/mfa/token/:token/devices", h.WithLimiter(h.getMFADevicesWithTokenHandle))
h.POST("/webapi/mfa/token/:token/authenticatechallenge", h.WithLimiter(h.createAuthenticateChallengeWithTokenHandle))
h.POST("/webapi/mfa/token/:token/registerchallenge", h.WithLimiter(h.createRegisterChallengeWithTokenHandle))
// MFA private endpoints.
h.GET("/webapi/mfa/devices", h.WithAuth(h.getMFADevicesHandle))
h.POST("/webapi/mfa/authenticatechallenge", h.WithAuth(h.createAuthenticateChallengeHandle))
h.POST("/webapi/mfa/devices", h.WithAuth(h.addMFADeviceHandle))
// Device Trust.
// Do not enforce bearer token for /webconfirm, it is called from outside the
// Web UI.
h.GET("/webapi/devices/webconfirm", h.WithSession(h.deviceWebConfirm))
// trusted clusters
h.POST("/webapi/trustedclusters/validate", h.WithUnauthenticatedLimiter(h.validateTrustedCluster))
// User Status (used by client to check if user session is valid)
h.GET("/webapi/user/status", h.WithAuth(h.getUserStatus))
// TODO(kimlisa): DELETE IN 20.0 along with the api path defined in `config.ts`
// Replaced by its v2 endpoint
h.GET("/webapi/roles", h.WithAuth(h.listRolesHandle))
// v2 introduces a query param for optionally including system roles
// in the list and optionally include returning the object version
// of resource (only supported for roles).
h.GET("/v2/webapi/roles", h.WithAuth(h.listRolesHandle))
h.POST("/webapi/roles", h.WithAuth(h.createRoleHandle))
h.GET("/webapi/roles/:name", h.WithAuth(h.getRole))
h.PUT("/webapi/roles/:name", h.WithAuth(h.updateRoleHandle))
h.DELETE("/webapi/roles/:name", h.WithAuth(h.deleteRole))
h.GET("/webapi/requestableroles", h.WithAuth(h.listRequestableRolesHandle))
h.GET("/webapi/presetroles", h.WithUnauthenticatedHighLimiter(h.getPresetRoles))
h.GET("/webapi/github", h.WithAuth(h.getGithubConnectorsHandle))
h.POST("/webapi/github", h.WithAuth(h.createGithubConnectorHandle))
// The extra "connector" in the path is to avoid a wildcard conflict with the github handlers used
// during the login flow ("github/login/web" and "github/callback").
h.GET("/webapi/github/connector/:name", h.WithAuth(h.getGithubConnectorHandle))
h.PUT("/webapi/github/:name", h.WithAuth(h.updateGithubConnectorHandle))
h.DELETE("/webapi/github/:name", h.WithAuth(h.deleteGithubConnector))
// Sets the default connector in the auth preference.
h.PUT("/webapi/authconnector/default", h.WithAuth(h.setDefaultConnectorHandle))
// Returns auth connectors that match a given username.
h.POST("/webapi/authconnectors", h.WithLimiter(h.getUserMatchedAuthConnectors))
h.GET("/webapi/trustedcluster", h.WithAuth(h.getTrustedClustersHandle))
h.POST("/webapi/trustedcluster", h.WithAuth(h.upsertTrustedClusterHandle))
h.PUT("/webapi/trustedcluster/:name", h.WithAuth(h.upsertTrustedClusterHandle))
h.DELETE("/webapi/trustedcluster/:name", h.WithAuth(h.deleteTrustedCluster))
h.GET("/webapi/apps/:fqdnHint", h.WithAuth(h.getAppDetails))
h.GET("/webapi/apps/:fqdnHint/:clusterName/:publicAddr", h.WithAuth(h.getAppDetails))
h.POST("/webapi/yaml/parse/:kind", h.WithAuth(h.yamlParse))
h.POST("/webapi/yaml/stringify/:kind", h.WithAuth(h.yamlStringify))
// Desktop access endpoints.
h.GET("/webapi/sites/:site/desktops", h.WithClusterAuth(h.clusterDesktopsGet))
h.GET("/webapi/sites/:site/desktopservices", h.WithClusterAuth(h.clusterDesktopServicesGet))
h.GET("/webapi/sites/:site/desktops/:desktopName", h.WithClusterAuth(h.getDesktopHandle))
// GET /webapi/sites/:site/desktops/:desktopName/connect?username=<username>&width=<width>&height=<height>
h.GET("/webapi/sites/:site/desktops/:desktopName/connect/ws", h.WithClusterAuthWebSocket(h.desktopConnectHandle, WithSubprotocols(tdpb.ProtocolName)))
// GET /webapi/sites/:site/desktopplayback/:sid/ws
h.GET("/webapi/sites/:site/desktopplayback/:sid/ws", h.WithClusterAuthWebSocket(h.desktopPlaybackHandle))
// GET /webapi/sites/:site/linuxdesktops/:desktopName/connect/ws?username=<username>&width=<width>&height=<height>
h.GET("/webapi/sites/:site/linuxdesktops/:desktopName/connect/ws", h.WithClusterAuthWebSocket(h.linuxDesktopConnectHandle, WithSubprotocols(tdpb.ProtocolName)))
h.GET("/webapi/sites/:site/desktops/:desktopName/active", h.WithClusterAuth(h.desktopIsActive))
// GET a Connection Diagnostics by its name
h.GET("/webapi/sites/:site/diagnostics/connections/:connectionid", h.WithClusterAuth(h.getConnectionDiagnostic))
// Diagnose a Connection
h.POST("/webapi/sites/:site/diagnostics/connections", h.WithClusterAuth(h.diagnoseConnection))
// Integrations CRUD
h.GET("/webapi/sites/:site/integrations", h.WithClusterAuth(h.integrationsList))
h.POST("/webapi/sites/:site/integrations", h.WithClusterAuth(h.integrationsCreate))
h.GET("/webapi/sites/:site/integrations/:name", h.WithClusterAuth(h.integrationsGet))
h.PUT("/webapi/sites/:site/integrations/:name", h.WithClusterAuth(h.integrationsUpdate))
h.GET("/webapi/sites/:site/integrations/:name/stats", h.WithClusterAuth(h.integrationStats))
h.GET("/webapi/sites/:site/integrations/:name/discoveryrules", h.WithClusterAuth(h.integrationDiscoveryRules))
h.GET("/webapi/sites/:site/integrations/:name/ca", h.WithClusterAuth(h.integrationsExportCA))
// TODO(kimlisa): DELETE IN 19.0 - Replaced by /v2 equivalent endpoint
h.DELETE("/webapi/sites/:site/integrations/:name_or_subkind", h.WithClusterAuth(h.integrationsDelete))
h.DELETE("/v2/webapi/sites/:site/integrations/:name_or_subkind", h.WithClusterAuth(h.integrationsDelete))
// GET the Microsoft Teams plugin app.zip file.
h.GET("/webapi/sites/:site/plugins/:plugin/files/msteams_app.zip", h.WithClusterAuth(h.integrationsMsTeamsAppZipGet))
// AWS OIDC Integration Actions
h.GET("/webapi/scripts/integrations/configure/awsoidc-idp.sh", h.WithLimiter(h.awsOIDCConfigureIdP))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/ping", h.WithClusterAuth(h.awsOIDCPing))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/databases", h.WithClusterAuth(h.awsOIDCListDatabases))
h.GET("/webapi/scripts/integrations/configure/listdatabases-iam.sh", h.WithLimiter(h.awsOIDCConfigureListDatabasesIAM))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/deployservice", h.WithClusterAuth(h.awsOIDCDeployService))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/deploydatabaseservices", h.WithClusterAuth(h.awsOIDCDeployDatabaseServices))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/listdeployeddatabaseservices", h.WithClusterAuth(h.awsOIDCListDeployedDatabaseService))
h.GET("/webapi/scripts/integrations/configure/deployservice-iam.sh", h.WithLimiter(h.awsOIDCConfigureDeployServiceIAM))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/eksclusters", h.WithClusterAuth(h.awsOIDCListEKSClusters))
// TODO(kimlisa): DELETE IN 19.0 - replaced by /v2/webapi/sites/:site/integrations/aws-oidc/:name/enrolleksclusters
// MUST delete with related code found in web/packages/teleport/src/services/integrations/integrations.ts(enrollEksClusters)
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/enrolleksclusters", h.WithClusterAuth(h.awsOIDCEnrollEKSClusters))
// v2 endpoint introduces "extraLabels" field.
h.POST("/v2/webapi/sites/:site/integrations/aws-oidc/:name/enrolleksclusters", h.WithClusterAuth(h.awsOIDCEnrollEKSClusters))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/securitygroups", h.WithClusterAuth(h.awsOIDCListSecurityGroups))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/databasevpcs", h.WithClusterAuth(h.awsOIDCListDatabaseVPCs))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/subnets", h.WithClusterAuth(h.awsOIDCListSubnets))
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/requireddatabasesvpcs", h.WithClusterAuth(h.awsOIDCRequiredDatabasesVPCS))
h.GET("/webapi/scripts/integrations/configure/eks-iam.sh", h.WithLimiter(h.awsOIDCConfigureEKSIAM))
h.GET("/webapi/scripts/integrations/configure/access-graph-cloud-sync-iam.sh", h.WithLimiter(h.accessGraphCloudSyncOIDC))
h.GET("/webapi/scripts/integrations/configure/aws-oidc-bedrock.sh", h.WithLimiter(h.awsBedrockSummarizerOIDC))
h.GET("/webapi/scripts/integrations/configure/aws-app-access-iam.sh", h.WithLimiter(h.awsOIDCConfigureAWSAppAccessIAM))
// TODO(kimlisa): DELETE IN 19.0 - Replaced by /v2 equivalent endpoint
h.POST("/webapi/sites/:site/integrations/aws-oidc/:name/aws-app-access", h.WithClusterAuth(h.awsOIDCCreateAWSAppAccess))
// v2 endpoint introduces "labels" field
// MUST delete with related code found in web/packages/teleport/src/services/integrations/integrations.ts(createAwsAppAccess)
h.POST("/v2/webapi/sites/:site/integrations/aws-oidc/:name/aws-app-access", h.WithClusterAuth(h.awsOIDCCreateAWSAppAccess))
// The Integration DELETE endpoint already sets the expected named param after `/integrations/`
// It must be re-used here, otherwise the router will not start.
// See https://github.com/julienschmidt/httprouter/issues/364
h.DELETE("/webapi/sites/:site/integrations/:name_or_subkind/aws-app-access/:name", h.WithClusterAuth(h.awsOIDCDeleteAWSAppAccess))
h.GET("/webapi/scripts/integrations/configure/ec2-ssm-iam.sh", h.WithLimiter(h.awsOIDCConfigureEC2SSMIAM))
// AWS IAM Roles Anywhere Integration Actions
h.GET("/webapi/scripts/integrations/configure/awsra-trust-anchor.sh", h.WithLimiter(h.awsRolesAnywhereConfigureTrustAnchor))
h.POST("/webapi/sites/:site/integrations/aws-ra/:name/validate", h.WithClusterAuth(h.validateAWSRolesAnywhereIntegration))
h.POST("/webapi/sites/:site/integrations/aws-ra/:name/ping", h.WithClusterAuth(h.awsRolesAnywherePing))
h.POST("/webapi/sites/:site/integrations/aws-ra/:name/listprofiles", h.WithClusterAuth(h.awsRolesAnywhereListProfiles))
// SAML IDP integration endpoints
h.GET("/webapi/scripts/integrations/configure/gcp-workforce-saml.sh", h.WithLimiter(h.gcpWorkforceConfigScript))
// Okta integration endpoints.
h.GET(OktaJWKSWellknownURI, h.WithLimiter(h.jwksOkta))
// Azure OIDC integration endpoints
h.GET("/webapi/scripts/integrations/configure/azureoidc.sh", h.WithLimiter(h.azureOIDCConfigure))
// OIDC Integration specific endpoints:
// Unauthenticated access to OpenID Configuration - used for AWS OIDC IdP integration
h.GET("/.well-known/openid-configuration", h.WithLimiter(h.openidConfiguration))
h.GET(OIDCJWKWURI, h.WithLimiter(h.jwksOIDC))
h.GET("/webapi/thumbprint", h.WithLimiter(h.thumbprint))
// SPIFFE Federation Trust Bundle
h.GET("/webapi/spiffe/bundle.json", h.WithLimiter(h.getSPIFFEBundle))
h.GET("/workload-identity/jwt-jwks.json", h.WithLimiter(h.getSPIFFEJWKS))
h.GET("/workload-identity/.well-known/openid-configuration", h.WithLimiter(h.getSPIFFEOIDCDiscoveryDocument))
// DiscoveryConfig CRUD
h.GET("/webapi/sites/:site/discoveryconfig", h.WithClusterAuth(h.discoveryconfigList))
h.POST("/webapi/sites/:site/discoveryconfig", h.WithClusterAuth(h.discoveryconfigCreate))
h.GET("/webapi/sites/:site/discoveryconfig/:name", h.WithClusterAuth(h.discoveryconfigGet))
h.PUT("/webapi/sites/:site/discoveryconfig/:name", h.WithClusterAuth(h.discoveryconfigUpdate))
h.DELETE("/webapi/sites/:site/discoveryconfig/:name", h.WithClusterAuth(h.discoveryconfigDelete))
// User Tasks CRUD
// Listing Tasks by Integration: GET /webapi/sites/:site/usertask?integration=<integration-name>
h.GET("/webapi/sites/:site/usertask", h.WithClusterAuth(h.userTaskListByIntegration))
h.GET("/webapi/sites/:site/usertask/:name", h.WithClusterAuth(h.userTaskGet))
h.PUT("/webapi/sites/:site/usertask/:name/state", h.WithClusterAuth(h.userTaskStateUpdate))
// Connection upgrades.
h.GET("/webapi/connectionupgrade", httplib.MakeHandler(h.connectionUpgrade))
// create user events.
h.POST("/webapi/precapture", h.WithUnauthenticatedLimiter(h.createPreUserEventHandle))
// create authenticated user events.
h.POST("/webapi/capture", h.WithAuth(h.createUserEventHandle))
h.GET("/webapi/headless/:headless_authentication_id", h.WithAuth(h.getHeadless))
h.PUT("/webapi/headless/:headless_authentication_id", h.WithAuth(h.putHeadlessState))
h.PUT("/webapi/mfa/browser/:request_id", h.WithAuth(h.putBrowserMFA))
h.GET("/webapi/sites/:site/user-groups", h.WithClusterAuth(h.getUserGroups))
// Fetches the user's preferences
h.GET("/webapi/user/preferences", h.WithAuth(h.getUserPreferences))
// Updates the user's preferences
h.PUT("/webapi/user/preferences", h.WithAuth(h.updateUserPreferences))
// Fetches the user's cluster preferences.
h.GET("/webapi/user/preferences/:site", h.WithClusterAuth(h.getUserClusterPreferences))
// Updates the user's cluster preferences.
h.PUT("/webapi/user/preferences/:site", h.WithClusterAuth(h.updateUserClusterPreferences))
// Returns logins included in the Connect My Computer role of the user.
// Returns an empty list of logins if the user does not have a Connect My Computer role assigned.
h.GET("/webapi/connectmycomputer/logins", h.WithAuth(h.connectMyComputerLoginsList))
// Implements the agent version server.
// Channel can contain "/", hence the use of a catch-all parameter
h.GET("/webapi/automaticupgrades/channel/*request", h.WithUnauthenticatedHighLimiter(h.automaticUpgrades109))
// Managed updates
h.GET("/webapi/managedupdates", h.WithAuth(h.getManagedUpdatesDetails))
h.POST("/webapi/managedupdates/groups/:groupName/start", h.WithAuth(h.startGroupUpdate))
h.POST("/webapi/managedupdates/groups/:groupName/done", h.WithAuth(h.markGroupDone))
h.POST("/webapi/managedupdates/groups/:groupName/rollback", h.WithAuth(h.rollbackGroup))
// GET Machine ID bot by name
h.GET("/webapi/sites/:site/machine-id/bot/:name", h.WithClusterAuth(h.getBot))
// GET Machine ID bots
h.GET("/webapi/sites/:site/machine-id/bot", h.WithClusterAuth(h.listBots))
// Create Machine ID bots
h.POST("/webapi/sites/:site/machine-id/bot", h.WithClusterAuth(h.createBot))
// Create bot join tokens
h.POST("/webapi/sites/:site/machine-id/token", h.WithClusterAuth(h.createBotJoinToken))
// PUT Machine ID bot by name
// TODO(nicholasmarais1158) DELETE IN v20.0.0
// Replaced by `PUT /v2/webapi/sites/:site/machine-id/bot/:name` which allows editing more than just roles.
h.PUT("/webapi/sites/:site/machine-id/bot/:name", h.WithClusterAuth(h.updateBotV1))
// PUT Machine ID bot by name
// TODO(nicholasmarais1158) DELETE IN v20.0.0
// Replaced by `PUT /v3/webapi/sites/:site/machine-id/bot/:name` which allows editing description.
h.PUT("/v2/webapi/sites/:site/machine-id/bot/:name", h.WithClusterAuth(h.updateBotV2))
// PUT Machine ID bot by name
h.PUT("/v3/webapi/sites/:site/machine-id/bot/:name", h.WithClusterAuth(h.updateBotV3))
// Delete Machine ID bot
h.DELETE("/webapi/sites/:site/machine-id/bot/:name", h.WithClusterAuth(h.deleteBot))
// GET Machine ID instance for a bot by id
h.GET("/webapi/sites/:site/machine-id/bot/:name/bot-instance/:id", h.WithClusterAuth(h.getBotInstance))
// GET Machine ID bot instances (paged)
// TODO(nicholasmarais1158) DELETE IN v20.0.0
// Replaced by `GET /v2/webapi/sites/:site/machine-id/bot-instance`.
h.GET("/webapi/sites/:site/machine-id/bot-instance", h.WithClusterAuth(h.listBotInstances))
// GET Machine ID bot instances (paged)
h.GET("/v2/webapi/sites/:site/machine-id/bot-instance", h.WithClusterAuth(h.listBotInstancesV2))
// GET Machine ID bot instance metrics.
h.GET("/webapi/sites/:site/machine-id/bot-instance/metrics", h.WithClusterAuth(h.botInstanceMetrics))
// List workload identities
h.GET("/webapi/sites/:site/workload-identity", h.WithClusterAuth(h.listWorkloadIdentities))
// GET a paginated list of notifications for a user
h.GET("/webapi/sites/:site/notifications", h.WithClusterAuth(h.notificationsGet))
// Upsert the timestamp of the latest notification that the user has seen.
h.PUT("/webapi/sites/:site/lastseennotification", h.WithClusterAuth(h.notificationsUpsertLastSeenTimestamp))
// Upsert a notification state when to mark a notification as read or hide it.
h.PUT("/webapi/sites/:site/notificationstate", h.WithClusterAuth(h.notificationsUpsertNotificationState))
// Git servers
h.PUT("/webapi/sites/:site/gitservers", h.WithClusterAuth(h.gitServerCreateOrUpsert))
h.GET("/webapi/sites/:site/gitservers/:name", h.WithClusterAuth(h.gitServerGet))
h.DELETE("/webapi/sites/:site/gitservers/:name", h.WithClusterAuth(h.gitServerDelete))
h.GET("/webapi/sites/:site/sessionthumbnail/:session_id", h.WithClusterAuth(h.getSessionRecordingThumbnail))
h.GET("/webapi/sites/:site/sessionrecording/:session_id/metadata/ws", h.WithClusterAuthWebSocket(h.getSessionRecordingMetadata))
h.GET("/webapi/sites/:site/sessionrecording/:session_id/playback/ws", h.WithClusterAuthWebSocket(h.recordingPlaybackWS))
// MWI IaC Wizards
h.POST("/webapi/sites/:site/machine-id/wizards/ci-cd", h.WithClusterAuth(h.machineIDWizardGenerateIaC))
}
// GetProxyClient returns authenticated auth server client
func (h *Handler) GetProxyClient() authclient.ClientI {
return h.cfg.ProxyClient
}
// GetProxyClientCertificate returns the proxy client certificate.
func (h *Handler) GetProxyClientCertificate() (*tls.Certificate, error) {
if h.cfg.GetProxyClientCertificate == nil {
return nil, trace.BadParameter("GetProxyClientCertificate is not set")
}
tlsCert, err := h.cfg.GetProxyClientCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
return tlsCert, nil
}
// GetAccessPoint returns the caching access point.
func (h *Handler) GetAccessPoint() authclient.ProxyAccessPoint {
return h.cfg.AccessPoint
}
// Close closes associated session cache operations
func (h *Handler) Close() error {
return h.auth.Close()
}
type userStatusResponse struct {
RequiresDeviceTrust types.TrustedDeviceRequirement `json:"requiresDeviceTrust,omitempty"`
HasDeviceExtensions bool `json:"hasDeviceExtensions,omitempty"`
Message string `json:"message"` // Always set to "ok"
}
func (h *Handler) getUserStatus(w http.ResponseWriter, r *http.Request, _ httprouter.Params, c *SessionContext) (any, error) {
return userStatusResponse{
RequiresDeviceTrust: c.cfg.Session.GetTrustedDeviceRequirement(),
HasDeviceExtensions: c.cfg.Session.GetHasDeviceExtensions(),
Message: "ok",
}, nil
}
// handleGetUserOrResetToken has two handlers:
// - read user
// - return reset password token
// It has two because the expected route for reading a user overlaps with an already existing one
// Using `GET /webapi/users/:username` invalidates the `GET /webapi/users/password/token/:token` route
// An alternative would be using the resource's singular name `GET /webapi/user/:username` but it invalidates the `GET /webapi/user/status` route
// So, instead we'll use `GET /webapi/users/*wildcard`, parse the path/params and call the appropriate handler
func (h *Handler) handleGetUserOrResetToken(w http.ResponseWriter, r *http.Request, p httprouter.Params) {
// do we have multiple path fields or just one
relativePath := p.ByName("wildcard")
relativePath = strings.TrimPrefix(relativePath, "/") // relativePath might start with "/", removing it helps reasoning
pathFields := strings.Split(relativePath, "/")
params := httprouter.Params{}
handleFunc := httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
http.NotFound(w, r)
return nil, nil
})
// having one means we have an username
if len(pathFields) == 1 {
username, err := decodeURLPathParamField(pathFields[0])
if err != nil {
trace.WriteError(w, err)
return
}
params = httprouter.Params{httprouter.Param{
Key: "username",
Value: username,
}}
handleFunc = h.WithAuth(h.getUserHandle)
}
// if we have exactly 3 and they look like /password/token/:token
if len(pathFields) == 3 && pathFields[0] == "password" && pathFields[1] == "token" && pathFields[2] != "" {
token, err := decodeURLPathParamField(pathFields[2])
if err != nil {
trace.WriteError(w, err)
return
}
params = httprouter.Params{httprouter.Param{
Key: "token",
Value: token,
}}
handleFunc = httplib.MakeHandler(h.getResetPasswordTokenHandle)
}
handleFunc(w, r, params)
}
// decodeURLPathParamField URL-decodes a manually extracted path segment.
// Use this when path params are set manually, since httprouter.Params.ByName()
// only auto-decodes params captured by the router itself.
func decodeURLPathParamField(value string) (string, error) {
decoded, err := url.PathUnescape(value)
if err != nil {
return "", trace.BadParameter("failed to decode URL path segment: %v", err)
}
return decoded, nil
}
// getUserContext returns user context
//
// GET /webapi/sites/:site/context
func (h *Handler) getUserContext(w http.ResponseWriter, r *http.Request, p httprouter.Params, c *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
cn, err := h.cfg.AccessPoint.GetClusterName(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
if cn.GetClusterName() != cluster.GetName() {
return nil, trace.BadParameter("endpoint only implemented for root cluster")
}
accessChecker, err := c.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := c.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
user, err := clt.GetUser(r.Context(), c.GetUser(), false)
if err != nil {
return nil, trace.Wrap(err)
}
// The following section is similar to
// https://github.com/gravitational/teleport/blob/ea810d30d99f26e58a190edc5facfbe0c09ea5e5/lib/srv/desktop/windows_server.go#L757-L769
recConfig, err := c.cfg.UnsafeCachedAuthClient.GetSessionRecordingConfig(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
desktopRecordingEnabled := recConfig.GetMode() != types.RecordOff
features := h.GetClusterFeatures()
entitlement := modules.GetProtoEntitlement(&features, entitlements.AccessMonitoring)
// ensure entitlement is set & feature is configured
accessMonitoringEnabled := entitlement.Enabled && features.GetAccessMonitoringConfigured()
userContext, err := ui.NewUserContext(user, accessChecker.Roles(), features, desktopRecordingEnabled, accessMonitoringEnabled)
if err != nil {
return nil, trace.Wrap(err)
}
res, err := clt.GetAccessCapabilities(r.Context(), types.AccessCapabilitiesRequest{
RequestableRoles: true,
SuggestedReviewers: true,
})
if err != nil {
return nil, trace.Wrap(err)
}
userContext.AccessCapabilities = ui.AccessCapabilities{
RequestableRoles: res.RequestableRoles,
SuggestedReviewers: res.SuggestedReviewers,
RequireReason: res.RequireReason,
}
userContext.AllowedSearchAsRoles = accessChecker.GetAllowedSearchAsRoles()
userContext.Cluster, err = ui.GetClusterDetails(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
pingResp, err := clt.Ping(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
if pingResp.LicenseExpiry != nil && !pingResp.LicenseExpiry.IsZero() {
userContext.Cluster.LicenseExpiry = pingResp.LicenseExpiry
}
userContext.ConsumedAccessRequestID = c.cfg.Session.GetConsumedAccessRequestID()
if h.cfg.ScopesFeatures.Enabled {
assignments, err := stream.Collect(scopedutils.RangeScopedRoleAssignments(
r.Context(), clt.ScopedAccessServiceClient(), scopedaccessv1.ListScopedRoleAssignmentsRequest_builder{
// Note that we are using the AllCallerAssignments flag here rather
// than just looking for our assignments by username. This flag
// suppresses standard scope-pinning, which allows us to see
// assignments in parent/orthogonal scopes. Generally, scoped commands
// only show the subset of state subject to the currently pinned scope,
// but the purpose of this function is specifically to discover
// potential target scopes for logging in, so we want to see everything
// regardless of current scope.
AllCallerAssignments: true,
ScopeFilter: scopesv1.Filter_builder{
Mode: scopesv1.Mode_MODE_ALL,
}.Build(),
}.Build(),
))
if err != nil {
return nil, trace.Wrap(err)
}
var assignedScopes []string
for _, assignment := range assignments {
for _, subAssignment := range assignment.GetSpec().GetAssignments() {
assignedScopes = append(assignedScopes, subAssignment.GetScope())
}
}
// apply canonical sorting and deduplication
slices.SortFunc(assignedScopes, scopes.Sort)
assignedScopes = slices.CompactFunc(assignedScopes, func(a, b string) bool {
return scopes.Compare(a, b) == scopes.Equivalent
})
userContext.AvailableScopes = assignedScopes
userContext.Scope, err = scopeFromSessionContext(c)
if err != nil {
return nil, trace.Wrap(err)
}
}
return userContext, nil
}
func scopeFromSessionContext(c *SessionContext) (string, error) {
id, err := c.GetIdentity()
if err != nil {
return "", trace.Wrap(err)
}
if id == nil {
return "", nil
}
pin := id.ScopePin
if pin == nil {
return "", nil
}
return pin.GetScope(), nil
}
// PublicProxyAddr returns the publicly advertised proxy address
func (h *Handler) PublicProxyAddr() string {
return h.cfg.PublicProxyAddr
}
// AccessGraphAddr returns the TAG API address
func (h *Handler) AccessGraphAddr() utils.NetAddr {
return h.cfg.AccessGraphAddr
}
func localSettings(ctx context.Context, cap types.AuthPreference, m modules.Modules, logger *slog.Logger) (webclient.AuthenticationSettings, error) {
as := webclient.AuthenticationSettings{
Type: constants.Local,
SecondFactor: types.LegacySecondFactorFromSecondFactors(cap.GetSecondFactors()),
PreferredLocalMFA: cap.GetPreferredLocalMFA(),
AllowPasswordless: cap.GetAllowPasswordless(),
AllowHeadless: cap.GetAllowHeadless(),
Local: &webclient.LocalSettings{},
PrivateKeyPolicy: cap.GetPrivateKeyPolicy(),
PIVSlot: cap.GetPIVSlot(),
PIVPINCacheTTL: cap.GetPIVPINCacheTTL(),
DeviceTrust: deviceTrustSettings(cap, m),
SignatureAlgorithmSuite: cap.GetSignatureAlgorithmSuite(),
}
// Only copy the connector name if it's truly local and not a local fallback.
if cap.GetType() == constants.Local {
as.Local.Name = cap.GetConnectorName()
}
// U2F settings.
switch u2f, err := cap.GetU2F(); {
case err == nil:
as.U2F = &webclient.U2FSettings{AppID: u2f.AppID}
case !trace.IsNotFound(err):
logger.WarnContext(ctx, "Error reading U2F settings", "error", err)
}
// Webauthn settings.
switch webConfig, err := cap.GetWebauthn(); {
case err == nil:
as.Webauthn = &webclient.Webauthn{
RPID: webConfig.RPID,
}
case !trace.IsNotFound(err):
logger.WarnContext(ctx, "Error reading WebAuthn settings", "error", err)
}
return as, nil
}
func oidcSettings(connector types.OIDCConnector, cap types.AuthPreference, m modules.Modules) webclient.AuthenticationSettings {
return webclient.AuthenticationSettings{
Type: constants.OIDC,
OIDC: &webclient.OIDCSettings{
Name: connector.GetName(),
Display: connector.GetDisplay(),
IssuerURL: connector.GetIssuerURL(),
},
// Local fallback / MFA.
SecondFactor: types.LegacySecondFactorFromSecondFactors(cap.GetSecondFactors()),
PreferredLocalMFA: cap.GetPreferredLocalMFA(),
PrivateKeyPolicy: cap.GetPrivateKeyPolicy(),
PIVSlot: cap.GetPIVSlot(),
PIVPINCacheTTL: cap.GetPIVPINCacheTTL(),
DeviceTrust: deviceTrustSettings(cap, m),
SignatureAlgorithmSuite: cap.GetSignatureAlgorithmSuite(),
}
}
func samlSettings(connector types.SAMLConnector, cap types.AuthPreference, m modules.Modules) webclient.AuthenticationSettings {
return webclient.AuthenticationSettings{
Type: constants.SAML,
SAML: &webclient.SAMLSettings{
Name: connector.GetName(),
Display: connector.GetDisplay(),
SingleLogoutEnabled: connector.GetSingleLogoutURL() != "",
// Note that we get the connector's primary SSO field, not the MFA SSO field.
// These two values are often unique, but should have the same host prefix
// (e.g. https://dev-813354.oktapreview.com) in reasonable, functional setups.
SSO: connector.GetSSO(),
},
// Local fallback / MFA.
SecondFactor: types.LegacySecondFactorFromSecondFactors(cap.GetSecondFactors()),
PreferredLocalMFA: cap.GetPreferredLocalMFA(),
PrivateKeyPolicy: cap.GetPrivateKeyPolicy(),
PIVSlot: cap.GetPIVSlot(),
PIVPINCacheTTL: cap.GetPIVPINCacheTTL(),
DeviceTrust: deviceTrustSettings(cap, m),
SignatureAlgorithmSuite: cap.GetSignatureAlgorithmSuite(),
}
}
func githubSettings(connector types.GithubConnector, cap types.AuthPreference, m modules.Modules) webclient.AuthenticationSettings {
return webclient.AuthenticationSettings{
Type: constants.Github,
Github: &webclient.GithubSettings{
Name: connector.GetName(),
Display: connector.GetDisplay(),
EndpointURL: connector.GetEndpointURL(),
},
// Local fallback / MFA.
SecondFactor: types.LegacySecondFactorFromSecondFactors(cap.GetSecondFactors()),
PreferredLocalMFA: cap.GetPreferredLocalMFA(),
PrivateKeyPolicy: cap.GetPrivateKeyPolicy(),
PIVSlot: cap.GetPIVSlot(),
PIVPINCacheTTL: cap.GetPIVPINCacheTTL(),
DeviceTrust: deviceTrustSettings(cap, m),
SignatureAlgorithmSuite: cap.GetSignatureAlgorithmSuite(),
}
}
func deviceTrustSettings(cap types.AuthPreference, m modules.Modules) webclient.DeviceTrustSettings {
dt := cap.GetDeviceTrust()
return webclient.DeviceTrustSettings{
Disabled: deviceTrustDisabled(cap, m),
AutoEnroll: dt != nil && dt.AutoEnroll,
}
}
// deviceTrustDisabled is used to set its namesake field in
// [webclient.PingResponse.Auth].
func deviceTrustDisabled(cap types.AuthPreference, m modules.Modules) bool {
return dtconfig.GetEffectiveMode(cap.GetDeviceTrust(), m) == constants.DeviceTrustModeOff
}
func getAuthSettings(ctx context.Context, authClient authclient.ClientI, m modules.Modules, logger *slog.Logger) (webclient.AuthenticationSettings, error) {
authPreference, err := authClient.GetAuthPreference(ctx)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
var as webclient.AuthenticationSettings
switch authPreference.GetType() {
case constants.Local:
as, err = localSettings(ctx, authPreference, m, logger)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
case constants.OIDC:
if authPreference.GetConnectorName() != "" {
oidcConnector, err := authClient.GetOIDCConnector(ctx, authPreference.GetConnectorName(), false)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
as = oidcSettings(oidcConnector, authPreference, m)
} else {
// TODO(okraport): DELETE IN v21.0.0, remove GetOIDCConnectors
oidcConnectors, err := clientutils.CollectWithFallback(ctx,
func(ctx context.Context, limit int, start string) ([]types.OIDCConnector, string, error) {
return authClient.ListOIDCConnectors(ctx, limit, start, false)
},
func(ctx context.Context) ([]types.OIDCConnector, error) {
return authClient.GetOIDCConnectors(ctx, false)
},
)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
if len(oidcConnectors) == 0 {
return webclient.AuthenticationSettings{}, trace.BadParameter("no oidc connectors found")
}
as = oidcSettings(oidcConnectors[0], authPreference, m)
}
case constants.SAML:
if authPreference.GetConnectorName() != "" {
samlConnector, err := authClient.GetSAMLConnectorWithValidationOptions(ctx, authPreference.GetConnectorName(), false, types.SAMLConnectorValidationFollowURLs(false))
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
as = samlSettings(samlConnector, authPreference, m)
} else {
// TODO(okraport): DELETE IN v21.0.0, remove GetSAMLConnectorsWithValidationOptions
samlConnectors, err := clientutils.CollectWithFallback(ctx,
func(ctx context.Context, limit int, start string) ([]types.SAMLConnector, string, error) {
return authClient.ListSAMLConnectorsWithOptions(ctx, limit, start, false, types.SAMLConnectorValidationFollowURLs(false))
},
func(ctx context.Context) ([]types.SAMLConnector, error) {
return authClient.GetSAMLConnectorsWithValidationOptions(ctx, false, types.SAMLConnectorValidationFollowURLs(false))
},
)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
if len(samlConnectors) == 0 {
return webclient.AuthenticationSettings{}, trace.BadParameter("no saml connectors found")
}
as = samlSettings(samlConnectors[0], authPreference, m)
}
case constants.Github:
if authPreference.GetConnectorName() != "" {
githubConnector, err := authClient.GetGithubConnector(ctx, authPreference.GetConnectorName(), false)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
as = githubSettings(githubConnector, authPreference, m)
} else {
// TODO(okraport): DELETE IN v21.0.0, remove GetGithubConnectors
githubConnectors, err := clientutils.CollectWithFallback(ctx,
func(ctx context.Context, limit int, start string) ([]types.GithubConnector, string, error) {
return authClient.ListGithubConnectors(ctx, limit, start, false)
},
func(ctx context.Context) ([]types.GithubConnector, error) {
return authClient.GetGithubConnectors(ctx, false)
},
)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
if len(githubConnectors) == 0 {
return webclient.AuthenticationSettings{}, trace.BadParameter("no github connectors found")
}
as = githubSettings(githubConnectors[0], authPreference, m)
}
default:
return webclient.AuthenticationSettings{}, trace.BadParameter("unknown type %v", authPreference.GetType())
}
as.HasMessageOfTheDay = authPreference.GetMessageOfTheDay() != ""
pingResp, err := authClient.Ping(ctx)
if err != nil {
return webclient.AuthenticationSettings{}, trace.Wrap(err)
}
as.LoadAllCAs = pingResp.LoadAllCAs
as.DefaultSessionTTL = authPreference.GetDefaultSessionTTL()
return as, nil
}
// traces forwards spans from the web ui to the upstream collector configured for the proxy. If tracing is
// disabled then the forwarding is a noop.
func (h *Handler) traces(w http.ResponseWriter, r *http.Request, _ httprouter.Params, _ *SessionContext) (any, error) {
body, err := utils.ReadAtMost(r.Body, teleport.MaxHTTPResponseSize)
if err != nil {
h.logger.ErrorContext(r.Context(), "Failed to read traces request", "error", err)
w.WriteHeader(http.StatusBadRequest)
return nil, nil
}
if err := r.Body.Close(); err != nil {
h.logger.WarnContext(r.Context(), "Failed to close traces request body", "error", err)
}
var data tracepb.TracesData
if err := (protojson.UnmarshalOptions{DiscardUnknown: true}).Unmarshal(body, &data); err != nil {
h.logger.ErrorContext(r.Context(), "Failed to unmarshal traces request", "error", err)
w.WriteHeader(http.StatusBadRequest)
return nil, nil
}
if len(data.ResourceSpans) == 0 {
w.WriteHeader(http.StatusBadRequest)
return nil, nil
}
// Unmarshalling of TraceId, SpanId, and ParentSpanId might all yield incorrect values. The raw values from
// OpenTelemetry-js are hex encoded, but the unmarshal call above will decode them as base64.
// In order to ensure the ids are in the right format and won't be rejected by the upstream collector
// we attempt to convert them back into the base64 and then hex decode them.
for _, resourceSpan := range data.ResourceSpans {
for _, scopeSpan := range resourceSpan.ScopeSpans {
for _, span := range scopeSpan.Spans {
// attempt to convert the trace id to the right format
if tid, err := oteltrace.TraceIDFromHex(base64.StdEncoding.EncodeToString(span.TraceId)); err == nil {
span.TraceId = tid[:]
}
// attempt to convert the span id to the right format
if sid, err := oteltrace.SpanIDFromHex(base64.StdEncoding.EncodeToString(span.SpanId)); err == nil {
span.SpanId = sid[:]
}
// attempt to convert the parent span id to the right format
if len(span.ParentSpanId) > 0 {
if psid, err := oteltrace.SpanIDFromHex(base64.StdEncoding.EncodeToString(span.ParentSpanId)); err == nil {
span.ParentSpanId = psid[:]
}
}
}
}
}
go func() {
// Because the uploading happens in a goroutine we cannot use the request scoped context
// since it will more than likely get canceled prior to the traces being uploaded. Use
// a background context with a lenient timeout to allow for a large number of spans to complete.
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
if err := h.cfg.TraceClient.UploadTraces(ctx, data.ResourceSpans); err != nil {
h.logger.ErrorContext(ctx, "Failed to upload traces", "error", err)
}
}()
w.WriteHeader(http.StatusOK)
return nil, nil
}
func (h *Handler) ping(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var err error
authSettings, err := getAuthSettings(r.Context(), h.cfg.ProxyClient, h.cfg.Modules, h.logger)
if err != nil {
return nil, trace.Wrap(err)
}
proxyConfig, err := h.cfg.ProxySettings.GetProxySettings(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
pr, err := h.cfg.ProxyClient.Ping(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
group := r.URL.Query().Get(webclient.AgentUpdateGroupParameter)
updaterID := r.URL.Query().Get(webclient.AgentUpdateIDParameter)
authSettings.Scopes = scopes.ScopesStatusToString(pr.ScopesStatus)
return webclient.PingResponse{
Auth: authSettings,
Proxy: *proxyConfig,
ServerVersion: teleport.Version,
MinClientVersion: teleport.MinClientSemVer().String(),
ClusterName: h.auth.clusterName,
AutomaticUpgrades: pr.ServerFeatures.GetAutomaticUpgrades(),
AutoUpdate: h.automaticUpdateSettings184(r.Context(), group, updaterID),
Edition: h.cfg.Modules.BuildType(),
FIPS: h.cfg.Modules.IsFIPSBuild(),
}, nil
}
func (h *Handler) find(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
group := r.URL.Query().Get(webclient.AgentUpdateGroupParameter)
cacheKey := "find"
if group != "" {
cacheKey += "-" + group
}
// cache the generic answer to avoid doing work for each request
resp, err := utils.FnCacheGet(r.Context(), h.findEndpointCache, cacheKey, func(ctx context.Context) (*webclient.PingResponse, error) {
proxyConfig, err := h.cfg.ProxySettings.GetProxySettings(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
authPref, err := h.cfg.AccessPoint.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
return &webclient.PingResponse{
Proxy: *proxyConfig,
Auth: webclient.AuthenticationSettings{
SignatureAlgorithmSuite: authPref.GetSignatureAlgorithmSuite(),
},
ServerVersion: teleport.Version,
MinClientVersion: teleport.MinClientSemVer().String(),
ClusterName: h.auth.clusterName,
Edition: h.cfg.Modules.BuildType(),
FIPS: h.cfg.Modules.IsFIPSBuild(),
AutoUpdate: h.automaticUpdateSettings184(ctx, group, "" /* updater UUID */),
}, nil
})
if err != nil {
return nil, trace.Wrap(err)
}
// Now we modulate the autoupdate answer on a per-request basis.
// We don't want to cache one answer per updater UUID, so we take the
// cached result and just override what we must.
updaterID := r.URL.Query().Get(webclient.AgentUpdateIDParameter)
if updaterID != "" {
resp.AutoUpdate = h.automaticUpdateSettings184(r.Context(), group, updaterID)
}
return resp, nil
}
func (h *Handler) pingWithConnector(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
authClient := h.cfg.ProxyClient
connectorName := p.ByName("connector")
cap, err := authClient.GetAuthPreference(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
pingResp, err := authClient.Ping(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
loadAllCAs := pingResp.LoadAllCAs
proxyConfig, err := h.cfg.ProxySettings.GetProxySettings(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
response := &webclient.PingResponse{
Proxy: *proxyConfig,
ServerVersion: teleport.Version,
MinClientVersion: teleport.MinClientSemVer().String(),
ClusterName: h.auth.clusterName,
}
hasMessageOfTheDay := cap.GetMessageOfTheDay() != ""
if slices.Contains(constants.SystemConnectors, connectorName) {
response.Auth, err = localSettings(r.Context(), cap, h.cfg.Modules, h.logger)
if err != nil {
return nil, trace.Wrap(err)
}
response.Auth.HasMessageOfTheDay = hasMessageOfTheDay
response.Auth.LoadAllCAs = loadAllCAs
response.Auth.Local.Name = connectorName // echo connector queried by caller
return response, nil
}
// collectorNames stores a list of the registered collector names so that
// in the event that no connector has matched, the list can be returned.
var collectorNames []string
// first look for a oidc connector with that name
oidcConnectors, err := authClient.GetOIDCConnectors(r.Context(), false)
if err == nil {
for index, value := range oidcConnectors {
collectorNames = append(collectorNames, value.GetMetadata().Name)
if value.GetMetadata().Name == connectorName {
response.Auth = oidcSettings(oidcConnectors[index], cap, h.cfg.Modules)
response.Auth.HasMessageOfTheDay = hasMessageOfTheDay
response.Auth.LoadAllCAs = loadAllCAs
return response, nil
}
}
}
// if no oidc connector was found, look for a saml connector
samlConnectors, err := authClient.GetSAMLConnectorsWithValidationOptions(r.Context(), false, types.SAMLConnectorValidationFollowURLs(false))
if err == nil {
for index, value := range samlConnectors {
collectorNames = append(collectorNames, value.GetMetadata().Name)
if value.GetMetadata().Name == connectorName {
response.Auth = samlSettings(samlConnectors[index], cap, h.cfg.Modules)
response.Auth.HasMessageOfTheDay = hasMessageOfTheDay
response.Auth.LoadAllCAs = loadAllCAs
return response, nil
}
}
}
// look for github connector
githubConnectors, err := authClient.GetGithubConnectors(r.Context(), false)
if err == nil {
for index, value := range githubConnectors {
collectorNames = append(collectorNames, value.GetMetadata().Name)
if value.GetMetadata().Name == connectorName {
response.Auth = githubSettings(githubConnectors[index], cap, h.cfg.Modules)
response.Auth.HasMessageOfTheDay = hasMessageOfTheDay
response.Auth.LoadAllCAs = loadAllCAs
return response, nil
}
}
}
return nil,
trace.BadParameter(
"invalid connector name: %v; valid options: %s",
connectorName, strings.Join(collectorNames, ", "))
}
// getWebConfig returns configuration for the web application.
func (h *Handler) getWebConfig(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
w.Header().Set("Content-Type", "application/javascript")
clusterFeatures := h.GetClusterFeatures()
automaticUpgradesEnabled := clusterFeatures.GetAutomaticUpgrades()
var (
oidcConnectors []types.OIDCConnector
samlConnectors []types.SAMLConnector
githubConnectors []types.GithubConnector
cap types.AuthPreference
proxyConfig *webclient.ProxySettings
recCfg types.SessionRecordingConfig
clusterName types.ClusterName
uiConfig webclient.UIConfig
sessionSummarizerEnabled bool
automaticUpgradesTargetVersion string
accessGraphConfigSet bool
)
var wg sync.WaitGroup
wg.Go(func() {
var err error
oidcConnectors, err = h.cfg.ProxyClient.GetOIDCConnectors(r.Context(), false)
if err != nil {
h.logger.ErrorContext(r.Context(), "Cannot retrieve OIDC connectors", "error", err)
}
})
wg.Go(func() {
var err error
samlConnectors, err = h.cfg.ProxyClient.GetSAMLConnectorsWithValidationOptions(r.Context(), false, types.SAMLConnectorValidationFollowURLs(false))
if err != nil {
h.logger.ErrorContext(r.Context(), "Cannot retrieve SAML connectors", "error", err)
}
})
wg.Go(func() {
var err error
githubConnectors, err = h.cfg.ProxyClient.GetGithubConnectors(r.Context(), false)
if err != nil {
h.logger.ErrorContext(r.Context(), "Cannot retrieve GitHub connectors", "error", err)
}
})
wg.Go(func() {
var err error
cap, err = h.cfg.AccessPoint.GetAuthPreference(r.Context())
if err != nil {
h.logger.ErrorContext(r.Context(), "Cannot retrieve AuthPreferences", "error", err)
}
})
wg.Go(func() {
var err error
proxyConfig, err = h.cfg.ProxySettings.GetProxySettings(r.Context())
if err != nil {
h.logger.WarnContext(r.Context(), "Cannot retrieve ProxySettings, tunnel address won't be set in Web UI", "error", err)
}
})
wg.Go(func() {
var err error
recCfg, err = h.cfg.AccessPoint.GetSessionRecordingConfig(r.Context())
if err != nil {
h.logger.ErrorContext(r.Context(), "Cannot retrieve SessionRecordingConfig", "error", err)
}
})
wg.Go(func() {
isEnabledRes, err := h.cfg.ProxyClient.SummarizerServiceClient().IsEnabled(
r.Context(),
&summarizerv1.IsEnabledRequest{},
)
if err == nil {
sessionSummarizerEnabled = isEnabledRes.GetEnabled()
}
})
wg.Go(func() {
rsp, err := h.cfg.ProxyClient.GetClusterAccessGraphConfig(r.Context())
if err != nil && !trace.IsNotImplemented(err) {
h.logger.ErrorContext(r.Context(), "Cannot retrieve Access Graph config from auth server", "error", err)
}
accessGraphConfigSet = rsp.GetEnabled() && rsp.GetAddress() != ""
})
wg.Go(func() {
var err error
clusterName, err = h.cfg.AccessPoint.GetClusterName(r.Context())
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to query cluster name", "error", err)
}
})
wg.Go(func() {
uiConfig = h.getUIConfig(r.Context())
})
if automaticUpgradesEnabled {
wg.Go(func() {
const group, updaterUUID = "", ""
agentVersion, err := h.autoUpdateResolver.GetVersion(r.Context(), group, updaterUUID)
if err != nil {
h.logger.ErrorContext(r.Context(), "Cannot read autoupdate target version", "error", err)
} else {
// agentVersion doesn't have the leading "v" which is expected here.
automaticUpgradesTargetVersion = fmt.Sprintf("v%s", agentVersion)
}
})
}
wg.Wait()
authProviders := []webclient.WebConfigAuthProvider{}
// identifierFirstLoginEnabled is true if at least one auth connector has a defined `user_matchers` field.
var identifierFirstLoginEnabled bool
for _, item := range oidcConnectors {
if item.GetUserMatchers() != nil {
identifierFirstLoginEnabled = true
}
authProviders = append(authProviders, webclient.WebConfigAuthProvider{
Type: webclient.WebConfigAuthProviderOIDCType,
WebAPIURL: webclient.WebConfigAuthProviderOIDCURL,
Name: item.GetName(),
DisplayName: item.GetDisplay(),
})
}
for _, item := range samlConnectors {
if item.GetUserMatchers() != nil {
identifierFirstLoginEnabled = true
}
authProviders = append(authProviders, webclient.WebConfigAuthProvider{
Type: webclient.WebConfigAuthProviderSAMLType,
WebAPIURL: webclient.WebConfigAuthProviderSAMLURL,
Name: item.GetName(),
DisplayName: item.GetDisplay(),
})
}
for _, item := range githubConnectors {
if item.GetUserMatchers() != nil {
identifierFirstLoginEnabled = true
}
authProviders = append(authProviders, webclient.WebConfigAuthProvider{
Type: webclient.WebConfigAuthProviderGitHubType,
WebAPIURL: webclient.WebConfigAuthProviderGitHubURL,
Name: item.GetName(),
DisplayName: item.GetDisplay(),
})
}
// get auth type & second factor type
var authSettings webclient.WebConfigAuthSettings
if cap == nil {
authSettings = webclient.WebConfigAuthSettings{
Providers: authProviders,
SecondFactor: constants.SecondFactorOff,
LocalAuthEnabled: true,
AuthType: constants.Local,
}
} else {
authType := cap.GetType()
var localConnectorName string
var defaultConnectorName string
if authType == constants.Local {
localConnectorName = cap.GetConnectorName()
} else {
defaultConnectorName = cap.GetConnectorName()
}
authSettings = webclient.WebConfigAuthSettings{
Providers: authProviders,
SecondFactor: types.LegacySecondFactorFromSecondFactors(cap.GetSecondFactors()),
SecondFactors: cap.GetSecondFactors(),
LocalAuthEnabled: cap.GetAllowLocalAuth(),
AllowPasswordless: cap.GetAllowPasswordless(),
AuthType: authType,
DefaultConnectorName: defaultConnectorName,
PreferredLocalMFA: cap.GetPreferredLocalMFA(),
LocalConnectorName: localConnectorName,
PrivateKeyPolicy: cap.GetPrivateKeyPolicy(),
MOTD: cap.GetMessageOfTheDay(),
IdentifierFirstLoginEnabled: identifierFirstLoginEnabled,
}
}
// get tunnel address to display on cloud instances
tunnelPublicAddr := ""
if proxyConfig != nil && clusterFeatures.GetCloud() {
tunnelPublicAddr = proxyConfig.SSH.TunnelPublicAddr
}
// disable joining sessions if proxy session recording is enabled
canJoinSessions := true
if recCfg != nil {
canJoinSessions = !services.IsRecordAtProxy(recCfg.GetMode())
}
disableRoleVisualizer, _ := strconv.ParseBool(os.Getenv("TELEPORT_UNSTABLE_DISABLE_ROLE_VISUALIZER"))
webCfg := webclient.WebConfig{
Edition: h.cfg.Modules.BuildType(),
Auth: authSettings,
CanJoinSessions: canJoinSessions,
IsCloud: clusterFeatures.GetCloud(),
TunnelPublicAddress: tunnelPublicAddr,
RecoveryCodesEnabled: clusterFeatures.GetRecoveryCodes(),
UI: uiConfig,
IsPolicyRoleVisualizerEnabled: !disableRoleVisualizer,
IsDashboard: services.IsDashboard(clusterFeatures),
IsUsageBasedBilling: clusterFeatures.GetIsUsageBased(),
AutomaticUpgrades: automaticUpgradesEnabled,
AutomaticUpgradesTargetVersion: automaticUpgradesTargetVersion,
CustomTheme: clusterFeatures.GetCustomTheme(),
Questionnaire: clusterFeatures.GetQuestionnaire(),
IsStripeManaged: clusterFeatures.GetIsStripeManaged(),
PremiumSupport: clusterFeatures.GetSupportType() == proto.SupportType_SUPPORT_TYPE_PREMIUM,
PlayableDatabaseProtocols: player.SupportedDatabaseProtocols,
SessionSummarizerEnabled: sessionSummarizerEnabled,
IsPolicyEnabled: modules.GetProtoEntitlement(&clusterFeatures, entitlements.Policy).Enabled,
// if Entitlements are not present, GetWebCfgEntitlements will return a map of entitlement to {enabled:false}
// if Entitlements are present, GetWebCfgEntitlements will populate the fields appropriately
Entitlements: getWebCfgEntitlements(&clusterFeatures),
IdentitySecurity: webclient.IdentitySecurity{
IsClusterLicensed: modules.GetProtoEntitlement(&clusterFeatures, entitlements.AccessGraph).Enabled ||
modules.GetProtoEntitlement(&clusterFeatures, entitlements.ActivityCenter).Enabled ||
modules.GetProtoEntitlement(&clusterFeatures, entitlements.SessionSummaries).Enabled,
AccessGraphConfigSet: accessGraphConfigSet,
SessionSummarizationEnabled: sessionSummarizerEnabled,
},
BeamsUI: clusterFeatures.GetBeamsUI(),
ScopesEnabled: h.cfg.ScopesFeatures.Enabled,
}
if clusterName != nil {
webCfg.ProxyClusterName = clusterName.GetClusterName()
}
out, err := json.Marshal(webCfg)
if err != nil {
return nil, trace.Wrap(err)
}
fmt.Fprintf(w, "var GRV_CONFIG = %v;", string(out))
return nil, nil
}
// getUserMatchedAuthConnectorsReq is the request body for getting user-matched auth connectors.
type getUserMatchedAuthConnectorsReq struct {
// Username is the username to match against the auth connectors.
Username string `json:"username"`
}
// getUserMatchedAuthConnectorsResponse is the response body for getting user-matched auth connectors.
type getUserMatchedAuthConnectorsResponse struct {
Connectors []webclient.WebConfigAuthProvider `json:"connectors"`
}
// globMatch performs simple a simple glob-style match test on a string.
// - '*' matches zero or more characters.
// It returns true if a match is detected.
func globMatch(pattern, str string) (bool, error) {
pattern = utils.GlobToRegexp(pattern)
matched, err := regexp.MatchString(pattern, str)
return matched, trace.Wrap(err)
}
// userMatchesConnector is a helper function to check if a user matches any of a connector's user matchers.
func userMatchesConnector(username string, connector interface {
GetUserMatchers() []string
},
) (isMatch bool, isFallback bool, err error) {
matchers := connector.GetUserMatchers()
for _, pattern := range matchers {
matched, err := globMatch(pattern, username)
if err != nil {
return false, false, trace.Wrap(err)
}
if matched {
// If the pattern is exactly "*", it matches all users and is considered a fallback.
// This match will only be returned to the user if the user doesn't match any other connectors via more explicit matchers.
// We continue here since there could be other matchers for this connector that are more specific, and if any of them match, this should not be treated as a fallback.
if pattern == "*" {
isFallback = true
continue
}
return true, false, nil
}
}
if isFallback {
return true, true, nil
}
return false, false, nil
}
// getUserMatchedAuthConnectors returns auth connectors that match the given username.
func (h *Handler) getUserMatchedAuthConnectors(w http.ResponseWriter, r *http.Request, params httprouter.Params) (any, error) {
var req *getUserMatchedAuthConnectorsReq
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if req.Username != "" && len(req.Username) > teleport.MaxUsernameLength {
return nil, trace.BadParameter("username exceeds maximum length of %d characters", teleport.MaxUsernameLength)
}
githubConns, err := h.cfg.ProxyClient.GetGithubConnectors(r.Context(), false)
if err != nil {
return nil, trace.Wrap(err)
}
samlConns, err := h.cfg.ProxyClient.GetSAMLConnectorsWithValidationOptions(r.Context(), false, types.SAMLConnectorValidationFollowURLs(false))
if err != nil {
return nil, trace.Wrap(err)
}
oidcConns, err := h.cfg.ProxyClient.GetOIDCConnectors(r.Context(), false)
if err != nil {
return nil, trace.Wrap(err)
}
var matchedConnectors []webclient.WebConfigAuthProvider
var fallbackConnectors []webclient.WebConfigAuthProvider
// Match GitHub connectors
for _, conn := range githubConns {
if matches, isFallback, err := userMatchesConnector(req.Username, conn); err != nil {
return nil, trace.Wrap(err)
} else if matches {
provider := webclient.WebConfigAuthProvider{
Type: webclient.WebConfigAuthProviderGitHubType,
WebAPIURL: webclient.WebConfigAuthProviderGitHubURL,
Name: conn.GetName(),
DisplayName: conn.GetDisplay(),
}
if isFallback {
fallbackConnectors = append(fallbackConnectors, provider)
} else {
matchedConnectors = append(matchedConnectors, provider)
}
}
}
// Match SAML connectors
for _, conn := range samlConns {
if matches, isFallback, err := userMatchesConnector(req.Username, conn); err != nil {
return nil, trace.Wrap(err)
} else if matches {
provider := webclient.WebConfigAuthProvider{
Type: webclient.WebConfigAuthProviderSAMLType,
WebAPIURL: webclient.WebConfigAuthProviderSAMLURL,
Name: conn.GetName(),
DisplayName: conn.GetDisplay(),
}
if isFallback {
fallbackConnectors = append(fallbackConnectors, provider)
} else {
matchedConnectors = append(matchedConnectors, provider)
}
}
}
// Match OIDC connectors
for _, conn := range oidcConns {
if matches, isFallback, err := userMatchesConnector(req.Username, conn); err != nil {
return nil, trace.Wrap(err)
} else if matches {
provider := webclient.WebConfigAuthProvider{
Type: webclient.WebConfigAuthProviderOIDCType,
WebAPIURL: webclient.WebConfigAuthProviderOIDCURL,
Name: conn.GetName(),
DisplayName: conn.GetDisplay(),
}
if isFallback {
fallbackConnectors = append(fallbackConnectors, provider)
} else {
matchedConnectors = append(matchedConnectors, provider)
}
}
}
// Use specific matches if available, otherwise use fallback matches.
var connectors []webclient.WebConfigAuthProvider
if len(matchedConnectors) > 0 {
connectors = matchedConnectors
} else {
connectors = fallbackConnectors
}
return &getUserMatchedAuthConnectorsResponse{
Connectors: connectors,
}, nil
}
// GetWebCfgEntitlements converts a proto entitlement map into the Web UI
// representation, including the legacy Policy fallback.
func GetWebCfgEntitlements(protoEntitlements map[string]*proto.EntitlementInfo) map[string]webclient.EntitlementInfo {
return getWebCfgEntitlements(&proto.Features{Entitlements: protoEntitlements})
}
func getWebCfgEntitlements(features *proto.Features) map[string]webclient.EntitlementInfo {
all := entitlements.AllEntitlements
result := make(map[string]webclient.EntitlementInfo, len(all))
for _, e := range all {
al := modules.GetProtoEntitlement(features, e)
result[string(e)] = webclient.EntitlementInfo{
Enabled: al.Enabled,
Limit: al.Limit,
}
}
return result
}
type JWKSResponse struct {
// Keys is a list of public keys in JWK format.
Keys []jwt.JWK `json:"keys"`
}
// getUiConfig will first try to get an ui_config set in the cache and then
// return what was set by the file config. Returns nil if neither are set which
// is fine, as the web UI can set its own defaults.
func (h *Handler) getUIConfig(ctx context.Context) webclient.UIConfig {
if uiConfig, err := h.cfg.AccessPoint.GetUIConfig(ctx); err == nil && uiConfig != nil {
return webclient.UIConfig{
ScrollbackLines: int(uiConfig.GetScrollbackLines()),
ShowResources: uiConfig.GetShowResources(),
}
}
return h.cfg.UI
}
// jwks returns all public keys used to sign JWT tokens for this cluster.
func (h *Handler) wellKnownJWKS(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
return h.jwks(r.Context(), types.JWTSigner, true)
}
func (h *Handler) motd(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
authPrefs, err := h.cfg.ProxyClient.GetAuthPreference(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
return webclient.MotD{Text: authPrefs.GetMessageOfTheDay()}, nil
}
func (h *Handler) githubLoginWeb(w http.ResponseWriter, r *http.Request, p httprouter.Params) string {
logger := h.logger.With("auth", "github")
logger.DebugContext(r.Context(), "Web login start")
req, err := ParseSSORequestParams(r)
if err != nil {
logger.ErrorContext(r.Context(), "Failed to extract SSO parameters from request", "error", err)
return sso.LoginFailedRedirectURL
}
remoteAddr, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
logger.ErrorContext(r.Context(), "Failed to parse request remote address", "error", err)
return sso.LoginFailedRedirectURL
}
response, err := h.cfg.ProxyClient.CreateGithubAuthRequest(r.Context(), types.GithubAuthRequest{
CSRFToken: req.CSRFToken,
ConnectorID: req.ConnectorID,
CreateWebSession: true,
ClientRedirectURL: req.ClientRedirectURL,
ClientLoginIP: remoteAddr,
ClientUserAgent: r.UserAgent(),
Scope: req.Scope,
})
if err != nil {
logger.ErrorContext(r.Context(), "Error creating auth request", "error", err)
return sso.LoginFailedRedirectURL
}
return response.RedirectURL
}
func (h *Handler) githubLoginConsole(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
logger := h.logger.With("auth", "github")
logger.DebugContext(r.Context(), "Console login start")
req := new(client.SSOLoginConsoleReq)
if err := httplib.ReadResourceJSON(r, req); err != nil {
logger.ErrorContext(r.Context(), "Error reading json", "error", err)
return nil, trace.AccessDenied("%s", SSOLoginFailureMessage)
}
if err := req.CheckAndSetDefaults(); err != nil {
logger.ErrorContext(r.Context(), "Missing request parameters", "error", err)
return nil, trace.AccessDenied("%s", SSOLoginFailureMessage)
}
remoteAddr, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
logger.ErrorContext(r.Context(), "Failed to parse request remote address", "error", err)
return nil, trace.AccessDenied("%s", SSOLoginFailureMessage)
}
response, err := h.cfg.ProxyClient.CreateGithubAuthRequest(r.Context(), types.GithubAuthRequest{
ConnectorID: req.ConnectorID,
SshPublicKey: req.SSHPubKey,
TlsPublicKey: req.TLSPubKey,
SshAttestationStatement: req.SSHAttestationStatement.ToProto(),
TlsAttestationStatement: req.TLSAttestationStatement.ToProto(),
CertTTL: req.CertTTL,
ClientRedirectURL: req.RedirectURL,
Compatibility: req.Compatibility,
RouteToCluster: req.RouteToCluster,
KubernetesCluster: req.KubernetesCluster,
ClientLoginIP: remoteAddr,
Scope: req.Scope,
})
if err != nil {
logger.ErrorContext(r.Context(), "Failed to create GitHub auth request", "error", err)
if strings.Contains(err.Error(), auth.InvalidClientRedirectErrorMessage) {
return nil, trace.AccessDenied("%s", SSOLoginFailureInvalidRedirect)
}
return nil, trace.AccessDenied("%s", SSOLoginFailureMessage)
}
return &client.SSOLoginConsoleResponse{
RedirectURL: response.RedirectURL,
}, nil
}
func (h *Handler) githubCallback(w http.ResponseWriter, r *http.Request, p httprouter.Params) string {
logger := h.logger.With("auth", "github")
logger.DebugContext(r.Context(), "Callback start", "query", r.URL.Query())
response, err := h.cfg.ProxyClient.ValidateGithubAuthCallback(r.Context(), r.URL.Query())
if err != nil {
logger.ErrorContext(r.Context(), "Error while processing callback", "error", err)
// try to find the auth request, which bears the original client redirect URL.
// if found, use it to terminate the flow.
//
// this improves the UX by terminating the failed SSO flow immediately, rather than hoping for a timeout.
if requestID := r.URL.Query().Get("state"); requestID != "" {
if request, errGet := h.cfg.ProxyClient.GetGithubAuthRequest(r.Context(), requestID); errGet == nil && !request.CreateWebSession {
if redURL, errEnc := RedirectURLWithError(request.ClientRedirectURL, err); errEnc == nil {
return redURL.String()
}
}
}
if errors.Is(err, auth.ErrGithubNoRoles) {
return sso.LoginFailedUnauthorizedRedirectURL
}
return sso.LoginFailedBadCallbackRedirectURL
}
// if we created web session, set session cookie and redirect to original url
if response.Req.CreateWebSession {
logger.InfoContext(r.Context(), "Redirecting to web browser")
res := &SSOCallbackResponse{
CSRFToken: response.Req.CSRFToken,
Username: response.Username,
SessionName: response.Session.GetName(),
SessionExpiry: response.Session.Expiry(),
ClientRedirectURL: response.Req.ClientRedirectURL,
}
if err := SSOSetWebSessionAndRedirectURL(w, r, res, true); err != nil {
logger.ErrorContext(r.Context(), "Error setting web session.", "error", err)
return sso.LoginFailedRedirectURL
}
if dwt := response.Session.GetDeviceWebToken(); dwt != nil {
logger.DebugContext(r.Context(), "GitHub WebSession created with device web token")
// if a device web token is present, we must send the user to the device authorize page
// to upgrade the session.
redirectPath, err := BuildDeviceWebRedirectPath(dwt, res.ClientRedirectURL)
if err != nil {
logger.DebugContext(r.Context(), "Invalid device web token", "error", err)
}
return redirectPath
}
return res.ClientRedirectURL
}
logger.InfoContext(r.Context(), "Callback is redirecting to console login")
if len(response.Req.SSHPubKey)+len(response.Req.TLSPubKey) == 0 {
logger.ErrorContext(r.Context(), "Not a web or console login request")
return sso.LoginFailedRedirectURL
}
redirectURL, err := ConstructSSHResponse(AuthParams{
ClientRedirectURL: response.Req.ClientRedirectURL,
Username: response.Username,
Identity: response.Identity,
Session: response.Session,
Cert: response.Cert,
TLSCert: response.TLSCert,
HostSigners: response.HostSigners,
FIPS: h.cfg.FIPS,
ClientOptions: response.ClientOptions,
})
if err != nil {
logger.ErrorContext(r.Context(), "Error constructing ssh response", "error", err)
return sso.LoginFailedRedirectURL
}
return redirectURL.String()
}
// BuildDeviceWebRedirectPath constructs the redirect path for device web authorization.
// It takes a DeviceWebToken and an optional client redirect URL as input.
// The function formats a redirect path with the device ID and token from the provided DeviceWebToken.
// If the clientRedirectURL is provided, it's appended to the redirect path
// as a query parameter named "redirect_uri".
// Will always at least return "/web" path.
func BuildDeviceWebRedirectPath(dwt *types.DeviceWebToken, clientRedirectURL string) (string, error) {
const basePath = "/web"
if dwt == nil {
return basePath, trace.BadParameter("DeviceWebToken cannot be nil")
}
if dwt.Id == "" || dwt.Token == "" {
return basePath, trace.BadParameter("DeviceWebToken ID and Token cannot be empty")
}
redirectPath := fmt.Sprintf("/web/device/authorize/%s/%s", dwt.Id, dwt.Token)
if clientRedirectURL != "" {
redirectPath = fmt.Sprintf("%s?redirect_uri=%s", redirectPath, clientRedirectURL)
}
return redirectPath, nil
}
func (h *Handler) installer(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
httplib.SetScriptHeaders(w.Header())
installerName := p.ByName("name")
installer, err := h.auth.proxyClient.GetInstaller(r.Context(), installerName)
if err != nil {
return nil, trace.Wrap(err)
}
ping, err := h.auth.Ping(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
const group, agentUUD = "", ""
targetVersion, err := h.autoUpdateResolver.GetVersion(r.Context(), group, agentUUD)
if err != nil {
h.logger.WarnContext(r.Context(), "Error retrieving the target version", "error", err)
targetVersion = teleport.SemVer()
}
getWindowsCA := func() (string, error) {
authorities, err := client.ExportAllAuthorities(
r.Context(),
h.GetProxyClient(),
client.ExportAuthoritiesRequest{
AuthType: "windows",
},
)
if err != nil {
return "", trace.Wrap(err)
}
// Combine all certificates into a single PEM-encoded string and then base64
// encode it.
var buf bytes.Buffer
for _, a := range authorities {
pem.Encode(&buf, &pem.Block{Type: "CERTIFICATE", Bytes: a.Data})
}
return base64.StdEncoding.EncodeToString(buf.Bytes()), nil
}
instTmpl, err := texttemplate.New("").
Funcs(texttemplate.FuncMap{"getWindowsCA": getWindowsCA}).
Parse(installer.GetScript())
if err != nil {
return nil, trace.Wrap(err)
}
teleportPackage := types.PackageNameOSS
if h.cfg.Modules.BuildType() == modules.BuildEnterprise || h.cfg.Modules.Features().Cloud {
teleportPackage = types.PackageNameEnt
if h.cfg.FIPS {
teleportPackage = types.PackageNameEntFIPS
}
}
// By default, it uses the stable/v<majorVersion> channel.
repoChannel := fmt.Sprintf("stable/v%d", targetVersion.Major)
// If the updater must be installed, then change the repo to stable/cloud
// It must also install the version specified in
// https://updates.releases.teleport.dev/v1/stable/cloud/version
installUpdater := ping.ServerFeatures.AutomaticUpgrades && ping.ServerFeatures.Cloud
if installUpdater {
repoChannel = automaticupgrades.DefaultCloudChannelName
}
azureClientID := r.URL.Query().Get("azure-client-id")
// For windows auth package installer scripts, we need to know if the installer
// should restart the machine after installation. A restart is required before
// smartcard authentication can be used, but we give the user the option
// because they may not want to restart immediately.
restartAfterEnrollment := r.URL.Query().Get("restart-after-enrollment") == "true"
tmpl := installers.Template{
PublicProxyAddr: h.PublicProxyAddr(),
MajorVersion: shsprintf.EscapeDefaultContext(fmt.Sprintf("v%d", targetVersion.Major)),
TeleportPackage: teleportPackage,
RepoChannel: shsprintf.EscapeDefaultContext(repoChannel),
AutomaticUpgrades: strconv.FormatBool(installUpdater),
AzureClientID: shsprintf.EscapeDefaultContext(azureClientID),
AuthPackageVersion: strings.ReplaceAll(targetVersion.String(), "'", "''"),
RestartAfterEnrollment: restartAfterEnrollment,
WindowsInstallerDownloadFailure: int(installstatus.WindowsInstallerDownloadFailure),
WindowsInstallerExecutionFailure: int(installstatus.WindowsInstallerExecutionFailure),
WindowsInstallerStagingDirUnsafe: int(installstatus.WindowsInstallerStagingDirUnsafe),
WindowsInstallerChecksumMismatch: int(installstatus.WindowsInstallerChecksumMismatch),
}
var buf bytes.Buffer
if err := instTmpl.Execute(&buf, tmpl); err != nil {
return nil, trace.Wrap(err)
}
if _, err := io.Copy(w, &buf); err != nil {
h.logger.DebugContext(r.Context(), "Failed writing installer script response", "error", err)
}
return nil, nil
}
// AuthParams are used to construct redirect URL containing auth
// information back to tsh login
type AuthParams struct {
// Username is authenticated teleport username
Username string
// Identity contains validated OIDC identity
Identity types.ExternalIdentity
// Web session will be generated by auth server if requested in OIDCAuthRequest
Session types.WebSession
// Cert will be generated by certificate authority
Cert []byte
// TLSCert is PEM encoded TLS certificate
TLSCert []byte
// HostSigners is a list of signing host public keys
// trusted by proxy, used in console login
HostSigners []types.CertAuthority
// ClientRedirectURL is a URL to redirect client to
ClientRedirectURL string
// FIPS mode means Teleport started in a FedRAMP/FIPS compliant
// configuration.
FIPS bool
// MFAToken is an SSO MFA token.
MFAToken string
// ClientOptions contains some options that the cluster wants the client to
// use.
ClientOptions authclient.ClientOptions
}
// ConstructSSHResponse creates a special SSH response for SSH login method
// that encodes everything using the client's secret key
func ConstructSSHResponse(response AuthParams) (*url.URL, error) {
u, err := url.Parse(response.ClientRedirectURL)
if err != nil {
return nil, trace.Wrap(err)
}
consoleResponse := authclient.CLILoginResponse{
Username: response.Username,
Cert: response.Cert,
TLSCert: response.TLSCert,
HostSigners: authclient.AuthoritiesToTrustedCerts(response.HostSigners),
MFAToken: response.MFAToken,
ClientOptions: response.ClientOptions,
}
out, err := json.Marshal(consoleResponse)
if err != nil {
return nil, trace.Wrap(err)
}
if u.Path == sso.WebMFARedirect {
// Transform the web sso mfa redirectURL into a relative redirect, preserving
// query parameters while ignoring scheme, hostname, and other url parts.
q := u.Query()
q.Add("response", string(out))
return &url.URL{
Path: sso.WebMFARedirect,
RawQuery: q.Encode(),
}, nil
}
// Extract secret out of the request.
secretKey := u.Query().Get("secret_key")
if secretKey == "" {
return nil, trace.BadParameter("missing secret_key")
}
var ciphertext []byte
// AES-GCM based symmetric cipher.
key, err := secret.ParseKey([]byte(secretKey))
if err != nil {
return nil, trace.Wrap(err)
}
ciphertext, err = key.Seal(out)
if err != nil {
return nil, trace.Wrap(err)
}
// Place ciphertext into the redirect URL.
u.RawQuery = url.Values{"response": {string(ciphertext)}}.Encode()
return u, nil
}
// RedirectURLWithError adds an err query parameter to the given redirect URL with the
// given errReply message and returns the new URL. If the given URL cannot be parsed,
// an error is returned with a nil URL. It is used to return an error back to the
// original URL in an SSO callback when validation fails.
func RedirectURLWithError(clientRedirectURL string, errReply error) (*url.URL, error) {
u, err := url.Parse(clientRedirectURL)
if err != nil {
return nil, trace.Wrap(err)
}
values := u.Query()
values.Set("err", errReply.Error())
u.RawQuery = values.Encode()
return u, nil
}
// CreateSessionReq is a request to create session from username, password and
// second factor token.
type CreateSessionReq struct {
// User is the Teleport username.
User string `json:"user"`
// Pass is the password.
Pass string `json:"pass"`
// SecondFactorToken is the OTP.
SecondFactorToken string `json:"second_factor_token"`
// Scope is the scope for which this session is created. Empty means
// unscoped.
Scope string `json:"scope"`
}
// String returns text description of this response
func (r *CreateSessionResponse) String() string {
return fmt.Sprintf("WebSession(type=%v,token=%v,expires=%vs)",
r.TokenType, r.Token, r.TokenExpiresIn)
}
// CreateSessionResponse returns OAuth compabible data about
// access token: https://tools.ietf.org/html/rfc6749
type CreateSessionResponse struct {
// TokenType is token type (bearer)
TokenType string `json:"type"`
// Token value
Token string `json:"token"`
// TokenExpiresIn sets seconds before this token is not valid
TokenExpiresIn int `json:"expires_in"`
// SessionExpiresIn is the seconds before the session itself expires.
SessionExpiresIn int `json:"sessionExpiresIn,omitempty"`
// SessionExpires is when this session expires.
SessionExpires time.Time `json:"sessionExpires"`
// SessionInactiveTimeoutMS specifies how long in milliseconds
// a user WebUI session can be left idle before being logged out
// by the server. A zero value means there is no idle timeout set.
SessionInactiveTimeoutMS int `json:"sessionInactiveTimeout"`
// DeviceWebToken is the token used to perform on-behalf-of device
// authentication.
// If not nil it should be forwarded to Connect for the device authentication
// ceremony.
DeviceWebToken *types.DeviceWebToken `json:"deviceWebToken,omitempty"`
// TrustedDeviceRequirement calculated for the web session.
// Follows [types.TrustedDeviceRequirement].
TrustedDeviceRequirement int32 `json:"trustedDeviceRequirement,omitempty"`
}
func newSessionResponse(ctx context.Context, sctx *SessionContext) (*CreateSessionResponse, error) {
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
if accessChecker.AccessInfo().ScopePin == nil {
_, err = accessChecker.CheckLoginDuration(0)
if err != nil {
return nil, trace.Wrap(err)
}
} else {
checkerContext, err := sctx.GetUserScopedAccessCheckerContext(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
_, err = checkerContext.CertParams().GetSSHLoginsForTTL(ctx, 0)
if err != nil {
return nil, trace.Wrap(err)
}
}
token, err := sctx.getToken()
if err != nil {
return nil, trace.Wrap(err)
}
now := sctx.cfg.Parent.clock.Now()
sessionExpiryTime := sctx.cfg.Session.GetExpiryTime()
return &CreateSessionResponse{
TokenType: roundtrip.AuthBearer,
Token: token.GetName(),
TokenExpiresIn: int(token.Expiry().Sub(now) / time.Second),
SessionExpiresIn: int(sessionExpiryTime.Sub(now) / time.Second),
SessionExpires: sessionExpiryTime,
SessionInactiveTimeoutMS: int(sctx.cfg.Session.GetIdleTimeout().Milliseconds()),
DeviceWebToken: sctx.cfg.Session.GetDeviceWebToken(),
TrustedDeviceRequirement: int32(sctx.cfg.Session.GetTrustedDeviceRequirement()),
}, nil
}
// createWebSession creates a new web session based on user, pass and 2nd factor token
//
// POST /v1/webapi/sessions/web
//
// {"user": "alex", "pass": "abcdef123456", "second_factor_token": "token", "second_factor_type": "totp"}
//
// # Response
//
// {"type": "bearer", "token": "bearer token", "user": {"name": "alex", "allowed_logins": ["admin", "bob"]}, "expires_in": 20}
func (h *Handler) createWebSession(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req *CreateSessionReq
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
// get cluster preferences to see if we should login
// with password or password+otp
authClient := h.cfg.ProxyClient
cap, err := authClient.GetAuthPreference(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
clientMeta := clientMetaFromReq(r)
clientMeta.ProxyGroupID = h.cfg.ProxyGroupID
var webSession types.WebSession
switch {
case !cap.IsSecondFactorEnforced():
webSession, err = h.auth.AuthWithoutOTP(r.Context(), req.User, req.Pass, req.Scope, clientMeta)
case req.SecondFactorToken == "" && !cap.IsSecondFactorEnforced():
webSession, err = h.auth.AuthWithoutOTP(r.Context(), req.User, req.Pass, req.Scope, clientMeta)
case cap.IsSecondFactorTOTPAllowed():
webSession, err = h.auth.AuthWithOTP(r.Context(), req.User, req.Pass, req.SecondFactorToken, req.Scope, clientMeta)
default:
return nil, trace.AccessDenied("direct login with password+otp not supported by this cluster")
}
if err != nil {
h.logger.WarnContext(r.Context(), "Access attempt denied for user", "user", req.User, "error", err)
// Since checking for private key policy meant that they passed authn,
// return policy error as is to help direct user.
if keys.IsPrivateKeyPolicyError(err) {
return nil, trace.Wrap(err)
}
// Obscure all other errors.
return nil, trace.AccessDenied("invalid credentials")
}
if err := websession.SetCookie(w, req.User, webSession.GetName(), webSession.Expiry()); err != nil {
return nil, trace.Wrap(err)
}
ctx, err := h.auth.newSessionContextFromSession(r.Context(), webSession)
if err != nil {
h.logger.WarnContext(r.Context(), "Access attempt denied for user", "user", req.User, "error", err)
return nil, trace.AccessDenied("need auth")
}
res, err := newSessionResponse(r.Context(), ctx)
return res, trace.Wrap(err)
}
func clientMetaFromReq(r *http.Request) *authclient.ForwardedClientMetadata {
var maxTouchPoints int
// The frontend client sends Max-Touch-Points only to endpoints that lead to the Device Trust
// prompt in the Web UI.
rawMaxTouchPoints := r.Header.Get("Max-Touch-Points")
if rawMaxTouchPoints != "" {
if value, err := strconv.Atoi(rawMaxTouchPoints); err == nil {
maxTouchPoints = value
}
}
return &authclient.ForwardedClientMetadata{
UserAgent: r.UserAgent(),
RemoteAddr: r.RemoteAddr,
MaxTouchPoints: maxTouchPoints,
}
}
// deleteWebSession is called to sign out user from web, app and SAML IdP session.
//
// DELETE /v1/webapi/sessions/:sid
//
// Response:
//
// {"message": "ok"}
func (h *Handler) deleteWebSession(w http.ResponseWriter, r *http.Request, _ httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to retrieve user client, SAML single logout will be skipped for user",
"user", ctx.GetUser(),
"error", err,
)
}
var user types.User
// Only run this if we successfully retrieved the client.
if err == nil {
user, err = clt.GetUser(r.Context(), ctx.GetUser(), false)
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to retrieve user during logout, SAML single logout will be skipped for user",
"user", ctx.GetUser(),
"error", err,
)
}
}
if err := h.logout(r.Context(), w, ctx); err != nil {
return nil, trace.Wrap(err)
}
// If the user has SAML SLO (single logout) configured, return a redirect link to the SLO URL.
if user != nil && len(user.GetSAMLIdentities()) > 0 && user.GetSAMLIdentities()[0].SAMLSingleLogoutURL != "" {
// The WebUI will redirect the user to this URL to initiate the SAML SLO on the IdP side. This is safe because this URL
// is hard-coded in the auth connector and can't be modified by the end user. Additionally, the user's Teleport session has already
// been invalidated by this point so there is nothing to hijack.
return map[string]any{"samlSloUrl": user.GetSAMLIdentities()[0].SAMLSingleLogoutURL}, nil
}
return OK(), nil
}
func (h *Handler) logout(ctx context.Context, w http.ResponseWriter, sctx *SessionContext) error {
if err := sctx.Invalidate(ctx); err != nil {
h.logger.WarnContext(ctx, "Failed to invalidate sessions",
"user", sctx.GetUser(),
"error", err,
)
}
if err := h.auth.releaseResources(ctx, sctx.GetUser(), sctx.GetSessionID()); err != nil {
h.logger.DebugContext(ctx, "sessionCache: Failed to release web session",
"session_id", sctx.GetSessionID(),
"error", err,
)
}
clearSessionCookies(w)
return nil
}
// clearSessionCookies clears Web UI session and SAML session cookie.
func clearSessionCookies(w http.ResponseWriter) {
// Clear Web UI session cookie
websession.ClearCookie(w)
}
type renewSessionRequest struct {
// AccessRequestID is the id of an approved access request.
AccessRequestID string `json:"requestId"`
// Switchback indicates switching back to default roles when creating new session.
Switchback bool `json:"switchback"`
// ReloadUser is a flag to indicate if user needs to be refetched from the backend
// to apply new user changes e.g. user traits were updated.
ReloadUser bool `json:"reloadUser"`
}
// renewWebSession updates this existing session with a new session.
//
// Depending on request fields sent in for extension, the new session creation can vary depending on:
// - AccessRequestID (opt): appends roles approved from access request to currently assigned roles or,
// - Switchback (opt): roles stacked with assuming approved access requests, will revert to user's default roles
// - ReloadUser (opt): similar to default but updates user related data (e.g login traits) by retrieving it from the backend
// - default (none set): create new session with currently assigned roles
func (h *Handler) renewWebSession(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
// TODO(bl-nero): Fix session renewal for scoped sessions. This endpoint
// doesn't work for scoped sessions, but currently, the web UI doesn't even
// call it in such scenario.
req := renewSessionRequest{}
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if req.AccessRequestID != "" && req.Switchback || req.AccessRequestID != "" && req.ReloadUser || req.Switchback && req.ReloadUser {
return nil, trace.BadParameter("failed to renew session: only one field can be set")
}
newSession, err := ctx.extendWebSession(r.Context(), req)
if err != nil {
return nil, trace.Wrap(err)
}
newContext, err := h.auth.newSessionContextFromSession(r.Context(), newSession)
if err != nil {
return nil, trace.Wrap(err)
}
if err := websession.SetCookie(w, newSession.GetUser(), newSession.GetName(), newSession.Expiry()); err != nil {
return nil, trace.Wrap(err)
}
res, err := newSessionResponse(r.Context(), newContext)
return res, trace.Wrap(err)
}
type changeUserAuthenticationRequest struct {
// SecondFactorToken is the TOTP code.
SecondFactorToken string `json:"second_factor_token"`
// TokenID is the ID of a reset or invite token.
TokenID string `json:"token"`
// DeviceName is the name of new mfa or passwordless device.
DeviceName string `json:"deviceName"`
// Password is user password string converted to bytes.
Password []byte `json:"password"`
// WebauthnCreationResponse is the signed credential creation response.
WebauthnCreationResponse *wantypes.CredentialCreationResponse `json:"webauthnCreationResponse"`
}
func (h *Handler) changeUserAuthentication(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req changeUserAuthenticationRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
protoReq := &proto.ChangeUserAuthenticationRequest{
TokenID: req.TokenID,
NewPassword: req.Password,
NewDeviceName: req.DeviceName,
}
switch {
case req.WebauthnCreationResponse != nil:
protoReq.NewMFARegisterResponse = &proto.MFARegisterResponse{
Response: &proto.MFARegisterResponse_Webauthn{
Webauthn: wantypes.CredentialCreationResponseToProto(req.WebauthnCreationResponse),
},
}
case req.SecondFactorToken != "":
protoReq.NewMFARegisterResponse = &proto.MFARegisterResponse{Response: &proto.MFARegisterResponse_TOTP{
TOTP: &proto.TOTPRegisterResponse{Code: req.SecondFactorToken},
}}
}
remoteAddr, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return nil, trace.Wrap(err)
}
protoReq.LoginIP = remoteAddr
res, err := h.auth.proxyClient.ChangeUserAuthentication(r.Context(), protoReq)
if err != nil {
return nil, trace.Wrap(err)
}
if res.PrivateKeyPolicyEnabled {
if res.GetRecovery() == nil {
return &ui.ChangedUserAuthn{
PrivateKeyPolicyEnabled: res.PrivateKeyPolicyEnabled,
}, nil
}
return &ui.ChangedUserAuthn{
Recovery: ui.RecoveryCodes{
Codes: res.GetRecovery().GetCodes(),
Created: &res.GetRecovery().Created,
},
PrivateKeyPolicyEnabled: res.PrivateKeyPolicyEnabled,
}, nil
}
sess := res.WebSession
ctx, err := h.auth.newSessionContextFromSession(r.Context(), sess)
if err != nil {
return nil, trace.Wrap(err)
}
if err := websession.SetCookie(w, sess.GetUser(), sess.GetName(), sess.Expiry()); err != nil {
return nil, trace.Wrap(err)
}
// Checks for at least one valid login.
if _, err := newSessionResponse(r.Context(), ctx); err != nil {
return nil, trace.Wrap(err)
}
if res.GetRecovery() == nil {
return &ui.ChangedUserAuthn{}, nil
}
return &ui.ChangedUserAuthn{
Recovery: ui.RecoveryCodes{
Codes: res.GetRecovery().GetCodes(),
Created: &res.GetRecovery().Created,
},
}, nil
}
// createResetPasswordToken allows a UI user to reset a user's password.
// This handler is also required for after creating new users.
func (h *Handler) createResetPasswordToken(w http.ResponseWriter, r *http.Request, _ httprouter.Params, ctx *SessionContext) (any, error) {
var req authclient.CreateUserTokenRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
token, err := clt.CreateResetPasswordToken(r.Context(),
authclient.CreateUserTokenRequest{
Name: req.Name,
Type: req.Type,
})
if err != nil {
return nil, trace.Wrap(err)
}
return ui.ResetPasswordToken{
TokenID: token.GetName(),
Expiry: token.Expiry(),
User: token.GetUser(),
}, nil
}
func (h *Handler) getResetPasswordTokenHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
result, err := h.getResetPasswordToken(r.Context(), p.ByName("token"))
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to fetch a reset password token", "error", err)
// We hide the error from the remote user to avoid giving any hints.
return nil, trace.AccessDenied("bad or expired token")
}
return result, nil
}
func (h *Handler) getResetPasswordToken(ctx context.Context, tokenID string) (any, error) {
token, err := h.auth.proxyClient.GetResetPasswordToken(ctx, tokenID)
if err != nil {
return nil, trace.Wrap(err)
}
// CreateRegisterChallenge rotates TOTP secrets for a given tokenID.
// It is required to get called every time a user fetches 2nd-factor secrets during registration attempt.
// This ensures that an attacker that gains the ResetPasswordToken link can not view it,
// extract the OTP key from the QR code, then allow the user to signup with
// the same OTP token.
res, err := h.auth.proxyClient.CreateRegisterChallenge(ctx, &proto.CreateRegisterChallengeRequest{
TokenID: tokenID,
DeviceType: proto.DeviceType_DEVICE_TYPE_TOTP,
})
if err != nil {
return nil, trace.Wrap(err)
}
return ui.ResetPasswordToken{
TokenID: token.GetName(),
User: token.GetUser(),
QRCode: res.GetTOTP().GetQRCode(),
}, nil
}
// mfaLoginBegin is the first step in the MFA authentication ceremony, which
// may be completed either via mfaLoginFinish (SSH) or mfaLoginFinishSession
// (Web).
//
// POST /webapi/mfa/login/begin
//
// {"user": "alex", "pass": "abcdef123456"}
// {"passwordless": true}
// {"user": "alex", "pass": "abcdef123456", "BrowserMFATSHRedirectURL": "http://localhost:12345/callback?secret_key=X"}
//
// Successful response:
//
// {"webauthn_challenge": {...}, "totp_challenge": true}
// {"webauthn_challenge": {...}} // passwordless
func (h *Handler) mfaLoginBegin(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req *client.MFAChallengeRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
mfaReq := &proto.CreateAuthenticateChallengeRequest{}
if req.Passwordless {
mfaReq.Request = &proto.CreateAuthenticateChallengeRequest_Passwordless{
Passwordless: &proto.Passwordless{},
}
mfaReq.ChallengeExtensions = &mfav1.ChallengeExtensions{
Scope: mfav1.ChallengeScope_CHALLENGE_SCOPE_PASSWORDLESS_LOGIN,
}
} else {
mfaReq.Request = &proto.CreateAuthenticateChallengeRequest_UserCredentials{
UserCredentials: &proto.UserCredentials{
Username: req.User,
Password: []byte(req.Pass),
},
}
mfaReq.ChallengeExtensions = &mfav1.ChallengeExtensions{
Scope: mfav1.ChallengeScope_CHALLENGE_SCOPE_LOGIN,
}
mfaReq.BrowserMFATSHRedirectURL = req.BrowserMFATSHRedirectURL
}
mfaChallenge, err := h.auth.proxyClient.CreateAuthenticateChallenge(r.Context(), mfaReq)
if err != nil {
// Do not obfuscate config-related errors.
if errors.Is(err, types.ErrPasswordlessRequiresWebauthn) || errors.Is(err, types.ErrPasswordlessDisabledBySettings) {
return nil, trace.Wrap(err)
}
return nil, trace.AccessDenied("invalid credentials")
}
return makeAuthenticateChallenge(mfaChallenge, "" /*channelID*/), nil
}
// mfaLoginFinish completes the MFA login ceremony, returning a new SSH
// certificate if successful.
//
// POST /v1/mfa/login/finish
//
// { "user": "bob", "password": "pass", "pub_key": "key to sign", "ttl": 1000000000 } # password-only
// { "user": "bob", "webauthn_challenge_response": {...}, "pub_key": "key to sign", "ttl": 1000000000 } # mfa
//
// # Success response
//
// { "cert": "base64 encoded signed cert", "host_signers": [{"domain_name": "example.com", "checking_keys": ["base64 encoded public signing key"]}] }
func (h *Handler) mfaLoginFinish(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req *client.AuthenticateSSHUserRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
clientMeta := clientMetaFromReq(r)
clientMeta.ProxyGroupID = h.cfg.ProxyGroupID
cert, err := h.auth.AuthenticateSSHUser(r.Context(), *req, clientMeta)
if err != nil {
return nil, trace.Wrap(err)
}
return cert, nil
}
// mfaLoginFinishSession completes the MFA login ceremony, returning a new web
// session if successful.
//
// POST /webapi/mfa/login/finishsession
//
// {"user": "alex", "webauthn_challenge_response": {...}}
//
// Successful response:
//
// {"type": "bearer", "token": "bearer token", "user": {"name": "alex", "allowed_logins": ["admin", "bob"]}, "expires_in": 20}
func (h *Handler) mfaLoginFinishSession(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
req := &client.AuthenticateWebUserRequest{}
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clientMeta := clientMetaFromReq(r)
clientMeta.ProxyGroupID = h.cfg.ProxyGroupID
session, err := h.auth.AuthenticateWebUser(r.Context(), req, clientMeta)
switch {
// Since checking for private key policy meant that they passed authn,
// return policy error as is to help direct user.
case keys.IsPrivateKeyPolicyError(err):
return nil, trace.Wrap(err)
// Return a friendlier error if an SSO user tried to do passwordless.
case errors.Is(err, types.ErrPassswordlessLoginBySSOUser):
return nil, trace.Wrap(err)
// Return a friendlier error if the user has assigned a role that doesn't exist in the
// backend.
case errors.Is(err, types.ErrNonExistingRoleAssigned):
return nil, trace.Wrap(err)
// Obscure all other errors.
case err != nil:
// log the actual error.
h.logger.WarnContext(r.Context(), "Login attempt denied for user", "user", req.User, "error", err)
return nil, trace.AccessDenied("invalid credentials")
}
// Fetch user from session, user is empty for passwordless requests.
user := session.GetUser()
if err := websession.SetCookie(w, user, session.GetName(), session.Expiry()); err != nil {
return nil, trace.Wrap(err)
}
ctx, err := h.auth.newSessionContextFromSession(r.Context(), session)
if err != nil {
return nil, trace.AccessDenied("need auth")
}
return newSessionResponse(r.Context(), ctx)
}
// getClusters returns a list of cluster and its data.
//
// GET /v1/webapi/sites
//
// Successful response:
//
// {"sites": {"name": "localhost", "last_connected": "RFC3339 time", "status": "active"}}
func (h *Handler) getClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, c *SessionContext) (any, error) {
// Get a client to the Auth Server with the logged in users identity. The
// identity of the logged in user is used to fetch the list of nodes.
clt, err := c.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
remoteClusters, err := clt.GetRemoteClusters(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
clusterName, err := clt.GetClusterName(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
rc, err := types.NewRemoteCluster(clusterName.GetClusterName())
if err != nil {
return nil, trace.Wrap(err)
}
rc.SetLastHeartbeat(time.Now().UTC())
rc.SetConnectionStatus(teleport.RemoteClusterStatusOnline)
clusters := make([]types.RemoteCluster, 0, len(remoteClusters)+1)
clusters = append(clusters, rc)
clusters = append(clusters, remoteClusters...)
out, err := ui.NewClustersFromRemote(clusters)
if err != nil {
return nil, trace.Wrap(err)
}
return out, nil
}
type getClusterInfoResponse struct {
ui.Cluster
IsCloud bool `json:"isCloud"`
}
// getClusterInfo returns the information about the cluster in the :site param
func (h *Handler) getClusterInfo(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
clusterDetails, err := ui.GetClusterDetails(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
pingResp, err := clt.Ping(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
return getClusterInfoResponse{
Cluster: *clusterDetails,
IsCloud: pingResp.GetServerFeatures().Cloud,
}, nil
}
type getSiteNamespacesResponse struct {
Namespaces []types.Namespace `json:"namespaces"`
}
// getSiteNamespaces returns a list of namespaces for a given site
//
// GET /v1/webapi/sites/:site/namespaces
//
// Successful response:
//
// {"namespaces": [{..namespace resource...}]}
func (h *Handler) getSiteNamespaces(w http.ResponseWriter, r *http.Request, _ httprouter.Params, c *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
return getSiteNamespacesResponse{
Namespaces: []types.Namespace{types.DefaultNamespace()},
}, nil
}
func makeUnifiedResourceRequest(r *http.Request, scopePin *scopesv1.Pin) (*proto.ListUnifiedResourcesRequest, error) {
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
sortBy := types.GetSortByFromString(values.Get("sort"))
var kinds []string
for _, kind := range values["kinds"] {
if kind != "" {
kinds = append(kinds, kind)
}
}
// include KindSAMLIdPServiceProvider when requesting KindApp
if slices.Contains(kinds, types.KindApp) &&
!slices.Contains(kinds, types.KindSAMLIdPServiceProvider) {
kinds = append(kinds, types.KindSAMLIdPServiceProvider)
}
// set default kinds to be requested if none exist in the request
scopedIdentity := scopePin.GetScope() != ""
if len(kinds) == 0 {
if scopedIdentity {
// TODO(bl-nero): Support more resource kinds when they are supported for
// scoped sessions.
kinds = []string{types.KindNode}
} else {
kinds = []string{
types.KindApp,
types.KindDatabase,
types.KindNode,
types.KindWindowsDesktop,
types.KindLinuxDesktop,
types.KindKubernetesCluster,
types.KindSAMLIdPServiceProvider,
types.KindGitServer,
}
}
}
startKey := values.Get("startKey")
includeRequestable := values.Get("includedResourceMode") == IncludedResourceModeAll
useSearchAsRoles := values.Get("searchAsRoles") == "yes" || values.Get("includedResourceMode") == IncludedResourceModeRequestable
return &proto.ListUnifiedResourcesRequest{
Kinds: kinds,
Limit: limit,
StartKey: startKey,
SortBy: sortBy,
PinnedOnly: values.Get("pinnedOnly") == "true",
PredicateExpression: values.Get("query"),
SearchKeywords: client.ParseSearchKeywords(values.Get("search"), ' '),
UseSearchAsRoles: useSearchAsRoles,
IncludeLogins: !scopedIdentity,
IncludeRequestable: includeRequestable,
}, nil
}
// getUserGroupLookup is a generator to retrieve UserGroupLookup on first call and return it again in subsequent calls.
// If we encounter an error, we log it once and return an empty UserGroupLookup for the current and subsequent calls.
// The returned function is not thread safe.
func (h *Handler) getUserGroupLookup(ctx context.Context, clt apiclient.GetResourcesClient) func() map[string]types.UserGroup {
userGroupLookup := make(map[string]types.UserGroup)
var gotUserGroupLookup bool
return func() map[string]types.UserGroup {
if gotUserGroupLookup {
return userGroupLookup
}
userGroups, err := apiclient.GetAllResources[types.UserGroup](ctx, clt, &proto.ListResourcesRequest{
ResourceType: types.KindUserGroup,
Namespace: apidefaults.Namespace,
UseSearchAsRoles: true,
})
if err != nil {
h.logger.InfoContext(ctx, "Unable to fetch user groups while listing applications, unable to display associated user groups", "error", err)
}
for _, userGroup := range userGroups {
userGroupLookup[userGroup.GetName()] = userGroup
}
gotUserGroupLookup = true
return userGroupLookup
}
}
// clusterUnifiedResourcesGet returns a list of resources for a given cluster site. This includes all resources available to be displayed in the web ui
// such as Nodes, Apps, Desktops, etc etc
func (h *Handler) clusterUnifiedResourcesGet(w http.ResponseWriter, request *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(request.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
identity, err := sctx.GetIdentity()
if err != nil {
return nil, trace.Wrap(err)
}
req, err := makeUnifiedResourceRequest(request, identity.ScopePin)
if err != nil {
return nil, trace.Wrap(err)
}
page, next, err := apiclient.GetUnifiedResourcePage(request.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
getUserGroupLookup := h.getUserGroupLookup(request.Context(), clt)
clusterAuthProxyServerFeatures := componentfeatures.GetClusterAuthProxyServerFeatures(request.Context(), h.GetAccessPoint(), h.logger)
unifiedResources := make([]any, 0, len(page))
for _, enriched := range page {
switch r := enriched.ResourceWithLabels.(type) {
case types.Server:
switch enriched.GetKind() {
case types.KindNode:
principals, err := PrincipalsForUnifiedResource(PrincipalsForUnifiedResourceOpts{
Resource: enriched,
CertPrincipals: identity.Principals,
AccessChecker: accessChecker,
UseSearchAsRoles: req.UseSearchAsRoles,
IncludeRequestable: req.IncludeRequestable,
})
if err != nil {
return nil, trace.Wrap(err)
}
nodeComponentFeatures := componentfeatures.Intersect(r.GetComponentFeatures(), clusterAuthProxyServerFeatures)
unifiedResources = append(unifiedResources, ui.MakeServer(r, ui.MakeServerConfig{
ClusterName: cluster.GetName(),
Logins: principals.Logins,
RequiresRequest: enriched.RequiresRequest,
SupportedFeatures: nodeComponentFeatures,
}))
case types.KindGitServer:
unifiedResources = append(unifiedResources, ui.MakeGitServer(cluster.GetName(), r, enriched.RequiresRequest))
}
case types.DatabaseServer:
db := ui.MakeDatabaseFromDatabaseServer(r, accessChecker, h.cfg.DatabaseREPLRegistry, enriched.RequiresRequest)
unifiedResources = append(unifiedResources, db)
case types.AppServer:
principals, err := PrincipalsForUnifiedResource(PrincipalsForUnifiedResourceOpts{
Resource: enriched,
AccessChecker: accessChecker,
IncludeRequestable: req.IncludeRequestable,
UseSearchAsRoles: req.UseSearchAsRoles,
})
if err != nil {
return nil, trace.Wrap(err)
}
proxyDNSName := h.proxyDNSName()
if r.GetApp().GetUseAnyProxyPublicAddr() {
// let the current proxy user is connected to override the dns name
proxyDNSName = utils.FindMatchingProxyDNS(request.Host, h.proxyDNSNames())
}
// Compute end-to-end feature support for this app: only features that are supported by the AppServer *and*
// by all required cluster hops (Auth + Proxy), so clients can hide features that would fail somewhere
// along the request path.
appComponentFeatures := componentfeatures.Intersect(componentfeatures.GetEffectiveServerFeatures(r), clusterAuthProxyServerFeatures)
app := ui.MakeApp(r.GetApp(), ui.MakeAppsConfig{
LocalClusterName: h.auth.clusterName,
LocalProxyDNSName: proxyDNSName,
AppClusterName: cluster.GetName(),
AWSRoles: principals.AWSRoleARNs,
UserGroupLookup: getUserGroupLookup(),
Logger: h.logger,
RequiresRequest: enriched.RequiresRequest,
SupportedFeatures: appComponentFeatures,
})
unifiedResources = append(unifiedResources, app)
case types.SAMLIdPServiceProvider:
// SAMLIdPServiceProvider resources are shown as
// "apps" in the UI.
app := ui.MakeAppTypeFromSAMLApp(r, ui.MakeAppsConfig{
LocalClusterName: h.auth.clusterName,
LocalProxyDNSName: h.proxyDNSName(),
AppClusterName: cluster.GetName(),
RequiresRequest: enriched.RequiresRequest,
})
unifiedResources = append(unifiedResources, app)
case types.WindowsDesktop:
logins := enriched.Logins
if req.IncludeRequestable || req.UseSearchAsRoles {
var err error
logins, err = accessChecker.GetAllowedLoginsForResource(r)
if err != nil {
return nil, trace.Wrap(err)
}
}
unifiedResources = append(unifiedResources, ui.MakeWindowsDesktop(r, logins, enriched.RequiresRequest))
case types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop]:
logins := enriched.Logins
if req.IncludeRequestable || req.UseSearchAsRoles {
var err error
logins, err = accessChecker.GetAllowedLoginsForResource(enriched)
if err != nil {
return nil, trace.Wrap(err)
}
}
unifiedResources = append(unifiedResources, ui.MakeLinuxDesktop(r.UnwrapT(), logins, enriched.RequiresRequest))
case types.KubeCluster:
kube := ui.MakeKubeCluster(r, accessChecker, enriched.RequiresRequest)
unifiedResources = append(unifiedResources, kube)
case types.KubeServer:
kube := ui.MakeKubeCluster(r.GetCluster(), accessChecker, enriched.RequiresRequest)
targetHealth := r.GetTargetHealth()
if targetHealth != nil {
kube.TargetHealth = *targetHealth
}
unifiedResources = append(unifiedResources, kube)
default:
return nil, trace.Errorf("UI Resource has unknown type: %T", enriched)
}
}
resp := listResourcesGetResponse{
Items: unifiedResources,
StartKey: next,
}
return resp, nil
}
// clusterNodesGet returns a list of nodes for a given cluster site.
func (h *Handler) clusterNodesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
// Get a client to the Auth Server with the logged in user's identity. The
// identity of the logged in user is used to fetch the list of nodes.
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindNode)
if err != nil {
return nil, trace.Wrap(err)
}
req.IncludeLogins = true
page, err := apiclient.GetEnrichedResourcePage(r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
identity, err := sctx.GetIdentity()
if err != nil {
return nil, trace.Wrap(err)
}
uiServers := make([]ui.Server, 0, len(page.Resources))
for _, resource := range page.Resources {
server, ok := resource.ResourceWithLabels.(types.Server)
if !ok {
continue
}
logins, err := client.CalculateSSHLogins(identity.Principals, resource.Logins)
if err != nil {
return nil, trace.Wrap(err)
}
loginSet := set.New(logins...)
uiServers = append(uiServers, ui.MakeServer(server, ui.MakeServerConfig{
ClusterName: cluster.GetName(),
Logins: &ui.PrincipalSet{All: loginSet, Granted: loginSet},
RequiresRequest: false,
SupportedFeatures: nil,
}))
}
return listResourcesGetResponse{
Items: uiServers,
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
// iso8601MilliFormat is the time format of dates returned from the frontend using Date().
const iso8601MilliFormat = "2006-01-02T15:04:05.000Z0700"
// notificationsGet returns a paginated list of notifications for a user.
func (h *Handler) notificationsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
startKey := values.Get("startKey")
response, err := clt.NotificationServiceClient().ListNotifications(r.Context(), notificationsv1.ListNotificationsRequest_builder{
PageSize: limit,
PageToken: startKey,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
var uiNotifications []ui.Notification
for _, notification := range response.GetNotifications() {
uiNotif := ui.MakeNotification(notification)
uiNotifications = append(uiNotifications, uiNotif)
}
return GetNotificationsResponse{
Notifications: uiNotifications,
NextKey: response.GetNextPageToken(),
UserLastSeenNotification: response.GetUserLastSeenNotificationTimestamp().AsTime().Format(iso8601MilliFormat),
}, nil
}
type GetNotificationsResponse struct {
Notifications []ui.Notification `json:"notifications"`
NextKey string `json:"nextKey"`
UserLastSeenNotification string `json:"userLastSeenNotification"`
}
// notificationsUpsertLastSeenTimestamp upserts a user's last seen notification timestamp.
func (h *Handler) notificationsUpsertLastSeenTimestamp(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
var req *UpsertUserLastSeenNotificationRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
lastSeenTime, err := time.Parse(iso8601MilliFormat, req.Time)
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := clt.NotificationServiceClient().UpsertUserLastSeenNotification(r.Context(), notificationsv1.UpsertUserLastSeenNotificationRequest_builder{
Username: sctx.GetUser(),
UserLastSeenNotification: notificationsv1.UserLastSeenNotification_builder{
Status: notificationsv1.UserLastSeenNotificationStatus_builder{
LastSeenTime: timestamppb.New(lastSeenTime),
}.Build(),
}.Build(),
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
return &UpsertUserLastSeenNotificationRequest{
Time: resp.GetStatus().GetLastSeenTime().AsTime().Format(iso8601MilliFormat),
}, nil
}
type UpsertUserLastSeenNotificationRequest struct {
Time string `json:"time"`
}
// notificationsUpsertNotificationState upserts a user notification state.
func (h *Handler) notificationsUpsertNotificationState(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
var req *upsertUserNotificationStateRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
resp, err := clt.NotificationServiceClient().UpsertUserNotificationState(r.Context(), notificationsv1.UpsertUserNotificationStateRequest_builder{
Username: sctx.GetUser(),
UserNotificationState: notificationsv1.UserNotificationState_builder{
Spec: notificationsv1.UserNotificationStateSpec_builder{
NotificationId: req.NotificationId,
}.Build(),
Status: notificationsv1.UserNotificationStateStatus_builder{
NotificationState: req.NotificationState,
}.Build(),
}.Build(),
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
return &upsertUserNotificationStateRequest{
NotificationId: resp.GetSpec().GetNotificationId(),
NotificationState: resp.GetStatus().GetNotificationState(),
}, nil
}
type upsertUserNotificationStateRequest struct {
NotificationId string `json:"notificationId"`
NotificationState notificationsv1.NotificationState `json:"notificationState"`
}
type getLoginAlertsResponse struct {
Alerts []types.ClusterAlert `json:"alerts"`
}
// clusterLoginAlertsGet returns a list of on-login alerts for the user.
func (h *Handler) clusterLoginAlertsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
// Get a client to the Auth Server with the logged in user's identity. The
// identity of the logged in user is used to fetch the list of alerts.
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
alerts, err := clt.GetClusterAlerts(h.cfg.Context, types.GetClusterAlertsRequest{
Labels: map[string]string{
types.AlertOnLogin: "yes",
},
})
if err != nil {
return nil, trace.Wrap(err)
}
return getLoginAlertsResponse{
Alerts: alerts,
}, nil
}
func (h *Handler) getClusterLocks(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sessionCtx *SessionContext,
cluster reversetunnelclient.Cluster,
) (any, error) {
ctx := r.Context()
clt, err := sessionCtx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
locks, err := clientutils.CollectWithFallback(
ctx,
func(ctx context.Context, limit int, start string) ([]types.Lock, string, error) {
var noFilter *types.LockFilter
return clt.ListLocks(ctx, limit, start, noFilter)
},
func(ctx context.Context) ([]types.Lock, error) {
// TODO(okraport): DELETE IN v21
const inForceOnlyFalse = false
return clt.GetLocks(ctx, inForceOnlyFalse)
},
)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeLocks(locks), nil
}
type GetClusterLocksV2Response struct {
Locks []ui.Lock `json:"items"`
}
func (h *Handler) getClusterLocksV2(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sessionCtx *SessionContext,
cluster reversetunnelclient.Cluster,
) (any, error) {
ctx := r.Context()
clt, err := sessionCtx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
// TODO(okraport): DELETE IN v21
var targets []types.LockTarget
filter := &types.LockFilter{}
if r.URL.Query().Has("target") {
ts := r.URL.Query()["target"]
for _, ts := range ts {
parts := strings.Split(ts, "|")
if len(parts) != 2 {
return nil, trace.BadParameter("invalid target %q", ts)
}
var target types.LockTarget
switch parts[0] {
case "user":
target.User = parts[1]
case "role":
target.Role = parts[1]
case "login":
target.Login = parts[1]
case "mfa_device":
target.MFADevice = parts[1]
case "windows_desktop":
target.WindowsDesktop = parts[1]
case "linux_desktop":
target.LinuxDesktop = parts[1]
case "access_request":
target.AccessRequest = parts[1]
case "device":
target.Device = parts[1]
case "server_id":
target.ServerID = parts[1]
case "bot_instance_id":
target.BotInstanceID = parts[1]
case "join_token":
target.JoinToken = parts[1]
default:
return nil, trace.BadParameter("invalid target type %q", parts[0])
}
targets = append(targets, target)
filter.Targets = append(filter.Targets, &target)
}
}
inForceOnly := false
if r.URL.Query().Has("in_force_only") {
inForceOnly = r.URL.Query().Get("in_force_only") == "true"
filter.InForceOnly = inForceOnly
}
locks, err := clientutils.CollectWithFallback(
ctx,
func(ctx context.Context, limit int, start string) ([]types.Lock, string, error) {
return clt.ListLocks(ctx, limit, start, filter)
},
func(ctx context.Context) ([]types.Lock, error) {
// TODO(okraport): DELETE IN v21
return clt.GetLocks(ctx, inForceOnly, targets...)
},
)
if err != nil {
return nil, trace.Wrap(err)
}
return &GetClusterLocksV2Response{
Locks: ui.MakeLocks(locks),
}, nil
}
type createLockReq struct {
Targets types.LockTarget `json:"targets"`
Message string `json:"message"`
TTL string `json:"ttl"`
}
func (h *Handler) createClusterLock(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sessionCtx *SessionContext,
cluster reversetunnelclient.Cluster,
) (any, error) {
var req *createLockReq
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
ctx := r.Context()
clt, err := sessionCtx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
var ttlDuration time.Duration
if req.TTL != "" {
ttlDuration, err = time.ParseDuration(req.TTL)
if err != nil {
return nil, trace.Wrap(err)
}
}
var expires *time.Time
if ttlDuration != 0 {
t := time.Now().UTC().Add(ttlDuration)
expires = &t
}
lock, err := types.NewLock(uuid.New().String(), types.LockSpecV2{
Target: req.Targets,
Message: req.Message,
Expires: expires,
})
if err != nil {
return nil, trace.Wrap(err)
}
err = clt.UpsertLock(ctx, lock)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeLock(lock), nil
}
func (h *Handler) deleteClusterLock(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sessionCtx *SessionContext,
cluster reversetunnelclient.Cluster,
) (any, error) {
ctx := r.Context()
clt, err := sessionCtx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
err = clt.DeleteLock(ctx, p.ByName("uuid"))
if err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// SessionController restricts creation of sessions based on
// cluster session control configuration(locks, connection limits, etc).
type SessionController interface {
// AcquireSessionContext attempts to create a context for the session. If the session is
// not allowed due to session control an error is returned. The returned
// context is scoped to the session and will be canceled in the event that session
// controls terminate the session early.
AcquireSessionContext(ctx context.Context, sctx *SessionContext, login, localAddr, remoteAddr string) (context.Context, error)
}
// SessionControllerFunc type is an adapter to allow the use of
// ordinary functions a [SessionController]. If f is a function
// with the appropriate signature, SessionControllerFunc(f) is a
// SessionController that calls f.
type SessionControllerFunc func(ctx context.Context, sctx *SessionContext, login, localAddr, remoteAddr string) (context.Context, error)
// AcquireSessionContext calls f(ctx, sctx, localAddr, remoteAddr).
func (f SessionControllerFunc) AcquireSessionContext(ctx context.Context, sctx *SessionContext, login, localAddr, remoteAddr string) (context.Context, error) {
ctx, err := f(ctx, sctx, login, localAddr, remoteAddr)
return ctx, trace.Wrap(err)
}
// siteNodeConnect connect to the site node
//
// GET /v1/webapi/sites/:site/namespaces/:namespace/connect?access_token=bearer_token¶ms=<urlencoded json-structure>
//
// Due to the nature of websocket we can't POST parameters as is, so we have
// to add query parameters. The params query parameter is a URL-encoded JSON structure:
//
// {"server_id": "uuid", "login": "admin", "term": {"h": 120, "w": 100}, "sid": "123"}
//
// Successful response is a websocket stream that allows read write to the server
func (h *Handler) siteNodeConnect(
w http.ResponseWriter,
r *http.Request,
_ httprouter.Params,
sessionCtx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (any, error) {
q := r.URL.Query()
params := q.Get("params")
if params == "" {
return nil, trace.BadParameter("missing params")
}
var req TerminalRequest
if err := json.Unmarshal([]byte(params), &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sessionCtx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ctx, err := h.cfg.SessionControl.AcquireSessionContext(r.Context(), sessionCtx, req.Login, h.cfg.ProxyWebAddr.Addr, r.RemoteAddr)
if err != nil {
return nil, trace.Wrap(err)
}
var (
sessionData session.Session
displayLogin string
tracker types.SessionTracker
)
clusterName := cluster.GetName()
if req.JoinSessionID.IsZero() {
sessionData, err = h.generateSession(r.Context(), &req, clusterName, sessionCtx)
if err != nil {
h.logger.DebugContext(r.Context(), "Unable to generate new ssh session", "error", err)
return nil, trace.Wrap(err)
}
} else {
sessionData, tracker, err = h.fetchJoinSession(ctx, clt, &req, clusterName)
if err != nil {
return nil, trace.Wrap(err)
}
displayLogin = tracker.GetLogin()
}
// If the participantMode is not specified, and the user is the one who created the session,
// they should be in Peer mode. If not, default to Observer mode.
if req.ParticipantMode == "" {
if sessionData.Owner == sessionCtx.cfg.User {
req.ParticipantMode = types.SessionPeerMode
} else {
req.ParticipantMode = types.SessionObserverMode
}
}
h.logger.DebugContext(r.Context(), "New terminal request",
"server", req.Server,
"login", req.Login,
"websid", sessionCtx.GetSessionID(),
)
authAccessPoint, err := cluster.CachingAccessPoint()
if err != nil {
h.logger.DebugContext(r.Context(), "Unable to get auth access point", "error", err)
return nil, trace.Wrap(err)
}
dialTimeout := apidefaults.DefaultIOTimeout
keepAliveInterval := apidefaults.KeepAliveInterval()
if netConfig, err := authAccessPoint.GetClusterNetworkingConfig(ctx); err != nil {
h.logger.DebugContext(r.Context(), "Unable to fetch cluster networking config", "error", err)
} else {
dialTimeout = netConfig.GetSSHDialTimeout()
keepAliveInterval = netConfig.GetKeepAliveInterval()
}
// Try to use the keep alive interval from the request.
// When it's not set or below a second, use the cluster's keep alive interval.
if req.KeepAliveInterval >= time.Second {
keepAliveInterval = req.KeepAliveInterval
}
nw, err := cluster.NodeWatcher()
if err != nil {
return nil, trace.Wrap(err)
}
term, err := NewTerminal(ctx, TerminalHandlerConfig{
Logger: h.logger,
Term: req.Term,
SessionCtx: sessionCtx,
UserAuthClient: clt,
LocalAccessPoint: h.auth.accessPoint,
DisplayLogin: displayLogin,
SessionData: sessionData,
KeepAliveInterval: keepAliveInterval,
ProxyHostPort: h.ProxyHostPort(),
ProxyPublicAddr: h.PublicProxyAddr(),
InteractiveCommand: req.InteractiveCommand,
Router: h.cfg.Router,
TracerProvider: h.cfg.TracerProvider,
ParticipantMode: req.ParticipantMode,
PROXYSigner: h.cfg.PROXYSigner,
Tracker: tracker,
PresenceChecker: h.cfg.PresenceChecker,
WebsocketConn: ws,
SSHDialTimeout: dialTimeout,
FIPSBuild: h.cfg.Modules.IsFIPSBuild(),
HostNameResolver: func(serverID string) (string, error) {
matches, err := nw.CurrentResourcesWithFilter(r.Context(), func(n readonly.Server) bool {
return n.GetName() == serverID
})
if err != nil {
return "", trace.Wrap(err)
}
if len(matches) != 1 {
return "", trace.NotFound("unable to resolve hostname for server %s", serverID)
}
return matches[0].GetHostname(), nil
},
})
if err != nil {
h.logger.ErrorContext(r.Context(), "Unable to create terminal", "error", err)
return nil, trace.Wrap(err)
}
h.userConns.Add(1)
defer h.userConns.Add(-1)
// start the websocket session with a web-based terminal:
httplib.MakeTracingHandler(term, teleport.ComponentProxy).ServeHTTP(w, r)
return nil, nil
}
func (h *Handler) setDefaultConnectorHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
var req ui.SetDefaultAuthConnectorRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
authPref, err := clt.GetAuthPreference(r.Context())
if err != nil {
return nil, trace.Wrap(err, "failed to get auth preference")
}
authPref.SetConnectorName(req.Name)
authPref.SetType(req.Type)
_, err = clt.UpsertAuthPreference(r.Context(), authPref)
if err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
type podConnectParams struct {
// Term is the initial PTY size.
Term session.TerminalParams `json:"term"`
// SessionID is a Teleport session ID to join as.
SessionID session.ID `json:"sid"`
// ParticipantMode is the mode that determines what you can do when you join an active session.
ParticipantMode types.SessionParticipantMode `json:"mode"`
}
func (h *Handler) podConnect(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (any, error) {
q := r.URL.Query()
if q.Get("params") == "" {
return nil, trace.BadParameter("missing params")
}
var params podConnectParams
if err := json.Unmarshal([]byte(q.Get("params")), ¶ms); err != nil {
return nil, trace.Wrap(err)
}
// If a session is provided, then join an existing session
// instead of creating a new one.
if !params.SessionID.IsZero() {
return nil, trace.Wrap(h.joinKubernetesSession(
r.Context(),
params.SessionID.String(),
params.ParticipantMode,
sctx,
cluster,
ws,
))
}
// Wait for the user to supply the pod information.
execReq, err := readPodExecRequestFromWS(ws)
if err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || terminal.IsOKWebsocketCloseError(trace.Unwrap(err)) {
return nil, nil
}
var netError net.Error
if errors.As(trace.Unwrap(err), &netError) && netError.Timeout() {
return nil, trace.BadParameter("timed out waiting for pod exec request data on websocket connection")
}
return nil, trace.Wrap(err)
}
execReq.Term = params.Term
if err := execReq.Validate(); err != nil {
return nil, trace.Wrap(err)
}
sess := session.Session{
Kind: types.KubernetesSessionKind,
Login: "root",
ClusterName: cluster.GetName(),
KubernetesClusterName: execReq.KubeCluster,
ID: session.NewID(),
Created: h.clock.Now().UTC(),
LastActive: h.clock.Now().UTC(),
Namespace: apidefaults.Namespace,
Owner: sctx.GetUser(),
Command: execReq.Command,
}
h.logger.DebugContext(r.Context(), "New kube exec request",
"namespace", execReq.Namespace,
"pod", execReq.Pod,
"container", execReq.Container,
"sid", sess.ID,
"websid", sctx.GetSessionID(),
)
authAccessPoint, err := cluster.CachingAccessPoint()
if err != nil {
return nil, trace.Wrap(err)
}
netConfig, err := authAccessPoint.GetClusterNetworkingConfig(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
serverAddr, tlsServerName, err := h.getKubeExecClusterData(netConfig)
if err != nil {
return nil, trace.Wrap(err)
}
hostCA, err := h.auth.accessPoint.GetCertAuthority(r.Context(), types.CertAuthID{
Type: types.HostCA,
DomainName: h.auth.clusterName,
}, false)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ph := podExecHandler{
req: execReq,
sess: sess,
sctx: sctx,
teleportCluster: cluster.GetName(),
ws: ws,
keepAliveInterval: netConfig.GetKeepAliveInterval(),
logger: h.logger.With(teleport.ComponentKey, "pod"),
userClient: clt,
localCA: hostCA,
configServerAddr: serverAddr,
configTLSServerName: tlsServerName,
publicProxyAddr: h.cfg.PublicProxyAddr,
}
ph.ServeHTTP(w, r)
return nil, nil
}
// KubeExecDataWaitTimeout is how long server would wait for user to send pod exec data (namespace, pod name etc)
// on websocket connection, after user initiated the exec into pod flow.
const KubeExecDataWaitTimeout = defaults.HeadlessLoginTimeout
func readPodExecRequestFromWS(ws *websocket.Conn) (*PodExecRequest, error) {
err := ws.SetReadDeadline(time.Now().Add(KubeExecDataWaitTimeout))
if err != nil {
return nil, trace.Wrap(err, "failed to set read deadline for websocket connection")
}
messageType, bytes, err := ws.ReadMessage()
if err != nil {
return nil, trace.Wrap(err)
}
if err := ws.SetReadDeadline(time.Time{}); err != nil {
return nil, trace.Wrap(err, "failed to set read deadline for websocket connection")
}
if messageType != websocket.BinaryMessage {
return nil, trace.BadParameter("Expected binary message of type websocket.BinaryMessage, got %v", messageType)
}
var envelope terminal.Envelope
if err := gogoproto.Unmarshal(bytes, &envelope); err != nil {
return nil, trace.BadParameter("Failed to parse envelope: %v", err)
}
var req PodExecRequest
if err := json.Unmarshal([]byte(envelope.Payload), &req); err != nil {
return nil, trace.Wrap(err)
}
return &req, nil
}
func (h *Handler) getKubeExecClusterData(netConfig types.ClusterNetworkingConfig) (string, string, error) {
if netConfig.GetProxyListenerMode() == types.ProxyListenerMode_Separate {
return "https://" + h.kubeProxyHostPort(), "", nil
}
proxyAddr := createHostPort(h.cfg.ProxyWebAddr, defaults.HTTPListenPort)
host, port, err := utils.SplitHostPort(proxyAddr)
if err != nil {
return "", "", trace.Wrap(err, "failed to split proxy address %q", proxyAddr)
}
tlsServerName := client.GetKubeTLSServerName(host)
return "https://" + net.JoinHostPort(host, port), tlsServerName, nil
}
func (h *Handler) generateSession(ctx context.Context, req *TerminalRequest, clusterName string, scx *SessionContext) (session.Session, error) {
owner := scx.cfg.User
h.logger.InfoContext(ctx, "Generating new session", "cluster", clusterName)
host, port, err := serverHostPort(req.Server)
if err != nil {
return session.Session{}, trace.Wrap(err)
}
accessChecker, err := scx.GetUserAccessChecker()
if err != nil {
return session.Session{}, trace.Wrap(err)
}
policySets := accessChecker.SessionPolicySets()
accessEvaluator := moderation.NewSessionAccessEvaluator(policySets, types.SSHSessionKind, owner)
return session.Session{
Kind: types.SSHSessionKind,
Login: req.Login,
ServerID: host,
ClusterName: clusterName,
ServerHostname: host,
ServerHostPort: port,
Moderated: accessEvaluator.IsModerated(),
Created: time.Now().UTC(),
LastActive: time.Now().UTC(),
Namespace: apidefaults.Namespace,
Owner: owner,
}, nil
}
// fetchJoinSession fetches an active or pending SSH session by the SessionID passed in the TerminalRequest.
func (h *Handler) fetchJoinSession(ctx context.Context, clt authclient.ClientI, req *TerminalRequest, siteName string) (session.Session, types.SessionTracker, error) {
// Session joining is not supported in proxy recording mode
if recConfig, err := h.auth.accessPoint.GetSessionRecordingConfig(ctx); err != nil {
// If the user can't see the recording mode, just let them try joining below
if !trace.IsAccessDenied(err) {
return session.Session{}, nil, trace.Wrap(err)
}
} else if services.IsRecordAtProxy(recConfig.GetMode()) {
return session.Session{}, nil, trace.BadParameter("session joining is not supported in proxy recording mode. If you are a Teleport administrator, you can learn more about recording modes at: https://goteleport.com/docs/reference/architecture/session-recording")
}
sessionID, err := session.ParseID(req.JoinSessionID.String())
if err != nil {
return session.Session{}, nil, trace.Wrap(err)
}
h.logger.InfoContext(ctx, "Attempting to join existing session", "session_id", sessionID)
tracker, err := clt.GetSessionTracker(ctx, string(*sessionID))
if err != nil {
return session.Session{}, nil, trace.Wrap(err)
}
if types.IsOpenSSHNodeSubKind(tracker.GetTargetSubKind()) {
return session.Session{}, nil, trace.BadParameter("session joining is only supported for nodes which are Teleport agents, not OpenSSH nodes")
}
if tracker.GetSessionKind() != types.SSHSessionKind || tracker.GetState() == types.SessionState_SessionStateTerminated {
return session.Session{}, nil, trace.NotFound("SSH session %v not found", sessionID)
}
sessionData := trackerToLegacySession(tracker, siteName)
// When joining an existing session use the specially handled
// `SSHSessionJoinPrincipal` login instead of the provided login so that
// users are able to join sessions without having permissions to create
// new ones themselves for auditing purposes. Otherwise, the user would
// fail the SSH lib username validation step.
sessionData.Login = teleport.SSHSessionJoinPrincipal
return sessionData, tracker, nil
}
type siteSessionGenerateResponse struct {
Session session.Session `json:"session"`
}
type siteSessionsGetResponse struct {
Sessions []siteSessionsGetResponseSession `json:"sessions"`
}
type siteSessionsGetResponseSession struct {
session.Session
ParticipantModes []types.SessionParticipantMode `json:"participantModes"`
}
func trackerToLegacySession(tracker types.SessionTracker, clusterName string) session.Session {
participants := tracker.GetParticipants()
parties := make([]session.Party, 0, len(participants))
for _, participant := range participants {
parties = append(parties, session.Party{
ID: session.ID(participant.ID),
User: participant.User,
Cluster: participant.Cluster,
ServerID: tracker.GetAddress(),
LastActive: participant.LastActive,
// note: we don't populate the RemoteAddr field since it isn't used and we don't have an equivalent value
})
}
accessEvaluator := moderation.NewSessionAccessEvaluator(tracker.GetHostPolicySets(), types.SSHSessionKind, tracker.GetHostUser())
return session.Session{
Kind: tracker.GetSessionKind(),
ID: session.ID(tracker.GetSessionID()),
Namespace: apidefaults.Namespace,
Parties: parties,
TerminalParams: session.TerminalParams{
W: teleport.DefaultTerminalWidth,
H: teleport.DefaultTerminalHeight,
},
Login: tracker.GetLogin(),
Created: tracker.GetCreated(),
LastActive: tracker.GetLastActive(),
ServerID: tracker.GetAddress(),
ServerHostname: tracker.GetHostname(),
ServerAddr: tracker.GetAddress(),
ClusterName: clusterName,
KubernetesClusterName: tracker.GetKubeCluster(),
DesktopName: tracker.GetDesktopName(),
AppName: tracker.GetAppName(),
Moderated: accessEvaluator.IsModerated(),
DatabaseName: tracker.GetDatabaseName(),
Owner: tracker.GetHostUser(),
Command: strings.Join(tracker.GetCommand(), " "),
}
}
// clusterActiveAndPendingSessionsGet gets the list of active and pending sessions for a site.
//
// GET /v1/webapi/sites/:site/sessions
func (h *Handler) clusterActiveAndPendingSessionsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
trackers, err := clt.GetActiveSessionTrackers(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
userRoles, err := clt.GetCurrentUserRoles(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
accessContext := moderation.SessionAccessContext{
Username: sctx.GetUser(),
Roles: userRoles,
}
sessions := make([]siteSessionsGetResponseSession, 0, len(trackers))
for _, tracker := range trackers {
if tracker.GetState() != types.SessionState_SessionStateTerminated {
session := trackerToLegacySession(tracker, p.ByName("site"))
// Get the participant modes available to the user from their roles.
accessEvaluator := moderation.NewSessionAccessEvaluator(tracker.GetHostPolicySets(), types.SSHSessionKind, tracker.GetHostUser())
participantModes := accessEvaluator.CanJoin(accessContext)
sessions = append(sessions, siteSessionsGetResponseSession{Session: session, ParticipantModes: participantModes})
}
}
return siteSessionsGetResponse{Sessions: sessions}, nil
}
func toFieldsSlice(rawEvents []apievents.AuditEvent) ([]events.EventFields, error) {
el := make([]events.EventFields, 0, len(rawEvents))
for _, event := range rawEvents {
els, err := events.ToEventFields(event)
if err != nil {
return nil, trace.Wrap(err)
}
el = append(el, els)
}
return el, nil
}
// clusterSearchEventsV2 returns all audit log events matching the provided criteria
//
// GET /v2/webapi/sites/:site/events/search
//
// Query parameters:
//
// "from" : date range from, encoded as RFC3339
// "to" : date range to, encoded as RFC3339
// "limit" : optional maximum number of events to return on each fetch
// "startKey": resume events search from the last event received,
// empty string means start search from beginning
// "include" : optional comma-separated list of event names to return e.g.
// include=session.start,session.end, all are returned if empty
// "order": optional ordering of events. Can be either "asc" or "desc"
// for ascending and descending respectively.
// If no order is provided it defaults to descending.
// "search": optional search term to filter events by (case-insensitive substring match)
func (h *Handler) clusterSearchEventsV2(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
values := r.URL.Query()
var eventTypes []string
if include := values.Get("include"); include != "" {
eventTypes = strings.Split(include, ",")
}
search := values.Get("search")
searchEvents := func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error) {
return clt.SearchEvents(r.Context(), events.SearchEventsRequest{
From: from,
To: to,
EventTypes: eventTypes,
Limit: limit,
Order: order,
StartKey: startKey,
Search: search,
})
}
return clusterEventsList(r.Context(), sctx, cluster, r.URL.Query(), searchEvents)
}
// clusterSearchEvents returns all audit log events matching the provided criteria
//
// GET /v1/webapi/sites/:site/events/search
//
// Query parameters:
//
// "from" : date range from, encoded as RFC3339
// "to" : date range to, encoded as RFC3339
// "limit" : optional maximum number of events to return on each fetch
// "startKey": resume events search from the last event received,
// empty string means start search from beginning
// "include" : optional comma-separated list of event names to return e.g.
// include=session.start,session.end, all are returned if empty
// "order": optional ordering of events. Can be either "asc" or "desc"
// for ascending and descending respectively.
// If no order is provided it defaults to descending.
func (h *Handler) clusterSearchEvents(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
values := r.URL.Query()
var eventTypes []string
if include := values.Get("include"); include != "" {
eventTypes = strings.Split(include, ",")
}
searchEvents := func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error) {
return clt.SearchEvents(r.Context(), events.SearchEventsRequest{
From: from,
To: to,
EventTypes: eventTypes,
Limit: limit,
Order: order,
StartKey: startKey,
})
}
return clusterEventsList(r.Context(), sctx, cluster, r.URL.Query(), searchEvents)
}
// clusterSearchSessionEvents returns session events matching the criteria.
//
// GET /v1/webapi/sites/:site/sessions/search
//
// Query parameters:
//
// "from" : date range from, encoded as RFC3339
// "to" : date range to, encoded as RFC3339
// "limit" : optional maximum number of events to return on each fetch
// "startKey": resume events search from the last event received,
// empty string means start search from beginning
// "order": optional ordering of events. Can be either "asc" or "desc"
// for ascending and descending respectively.
// If no order is provided it defaults to descending.
func (h *Handler) clusterSearchSessionEvents(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
searchSessionEvents := func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error) {
return clt.SearchSessionEvents(r.Context(), events.SearchSessionEventsRequest{
From: from,
To: to,
Limit: limit,
Order: order,
StartKey: startKey,
})
}
return clusterEventsList(r.Context(), sctx, cluster, r.URL.Query(), searchSessionEvents)
}
// clusterEventsList returns a list of audit events obtained using the provided
// searchEvents method.
func clusterEventsList(ctx context.Context, sctx *SessionContext, cluster reversetunnelclient.Cluster, values url.Values, searchEvents func(clt authclient.ClientI, from, to time.Time, limit int, order types.EventOrder, startKey string) ([]apievents.AuditEvent, string, error)) (any, error) {
from, err := queryTime(values, "from", time.Now().UTC().AddDate(0, -1, 0))
if err != nil {
return nil, trace.Wrap(err)
}
to, err := queryTime(values, "to", time.Now().UTC())
if err != nil {
return nil, trace.Wrap(err)
}
limit, err := QueryLimit(values, "limit", defaults.EventsIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
order, err := queryOrder(values, "order", types.EventOrderDescending)
if err != nil {
return nil, trace.Wrap(err)
}
startKey := values.Get("startKey")
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
rawEvents, lastKey, err := searchEvents(clt, from, to, limit, order, startKey)
if err != nil {
return nil, trace.Wrap(err)
}
el, err := toFieldsSlice(rawEvents)
if err != nil {
return nil, trace.Wrap(err)
}
return eventsListGetResponse{Events: el, StartKey: lastKey}, nil
}
// queryTime parses the query string parameter with the specified name as a
// RFC3339 time and returns it.
//
// If there's no such parameter, specified default value is returned.
func queryTime(query url.Values, name string, def time.Time) (time.Time, error) {
str := query.Get(name)
if str == "" {
return def, nil
}
parsed, err := time.Parse(time.RFC3339, str)
if err != nil {
return time.Time{}, trace.BadParameter("failed to parse %v as RFC3339 time: %v", name, str)
}
return parsed, nil
}
// QueryLimit returns the limit parameter with the specified name from the
// query string.
//
// If there's no such parameter, specified default limit is returned.
func QueryLimit(query url.Values, name string, def int) (int, error) {
str := query.Get(name)
if str == "" {
return def, nil
}
limit, err := strconv.Atoi(str)
if err != nil {
return 0, trace.BadParameter("failed to parse %v as limit: %v", name, str)
}
return limit, nil
}
// queryLimitAsInt32 returns the limit parameter with the specified name from the
// query string. Similar to function 'queryLimit' except it returns as type int32.
//
// If there's no such parameter, specified default limit is returned.
func QueryLimitAsInt32(query url.Values, name string, def int32) (int32, error) {
str := query.Get(name)
if str == "" {
return def, nil
}
limit, err := strconv.ParseInt(str, 10, 32)
if err != nil {
return 0, trace.BadParameter("failed to parse %v as limit: %v", name, str)
}
return int32(limit), nil
}
// queryOrder returns the order parameter with the specified name from the
// query string or a default if the parameter is not provided.
func queryOrder(query url.Values, name string, def types.EventOrder) (types.EventOrder, error) {
value := strings.ToLower(query.Get(name))
switch value {
case "desc":
return types.EventOrderDescending, nil
case "asc":
return types.EventOrderAscending, nil
case "":
return def, nil
default:
return types.EventOrderAscending, trace.BadParameter("parameter %v is not a valid ordering", value)
}
}
type eventsListGetResponse struct {
// Events is list of events retrieved.
Events []events.EventFields `json:"events"`
// StartKey is the position to resume search events.
StartKey string `json:"startKey"`
}
// hostCredentials sends a registration token and metadata to the Auth Server
// and gets back SSH and TLS certificates.
func (h *Handler) hostCredentials(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req types.RegisterUsingTokenRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
authClient := h.cfg.ProxyClient
remoteAddr, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return nil, trace.Wrap(err)
}
req.RemoteAddr = remoteAddr
certs, err := authClient.RegisterUsingToken(r.Context(), &req)
if err != nil {
return nil, trace.Wrap(err)
}
return certs, nil
}
// headlessLogin is a web call that perform headless login based on a user's
// name, returning new login certs if successful.
//
// POST /v1/webapi/headless/login
//
// { "user": "bob", "pub_key": "key to sign", "ttl": 1000000000 }
//
// # Success response
//
// { "cert": "base64 encoded signed cert", "host_signers": [{"domain_name": "example.com", "checking_keys": ["base64 encoded public signing key"]}] }
func (h *Handler) headlessLogin(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req client.HeadlessLoginReq
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
authClient := h.cfg.ProxyClient
clientMeta := clientMetaFromReq(r)
clientMeta.ProxyGroupID = h.cfg.ProxyGroupID
authSSHUserReq := authclient.AuthenticateSSHRequest{
AuthenticateUserRequest: authclient.AuthenticateUserRequest{
Username: req.User,
SSHPublicKey: req.SSHPubKey,
TLSPublicKey: req.TLSPubKey,
ClientMetadata: clientMeta,
HeadlessAuthenticationID: req.HeadlessAuthenticationID,
},
CompatibilityMode: req.Compatibility,
TTL: req.TTL,
RouteToCluster: req.RouteToCluster,
KubernetesCluster: req.KubernetesCluster,
SSHAttestationStatement: req.SSHAttestationStatement,
TLSAttestationStatement: req.TLSAttestationStatement,
}
// We need to use the default callback timeout rather than the standard client timeout.
// However, authClient is shared across all Proxy->Auth requests, so we need to create
// a new client to avoid applying the callback timeout to other concurrent requests. To
// this end, we create a clone of the HTTP Client with the desired timeout instead.
httpClient, err := authClient.CloneHTTPClient(
authclient.ClientParamTimeout(defaults.HeadlessLoginTimeout),
authclient.ClientParamResponseHeaderTimeout(defaults.HeadlessLoginTimeout),
)
if err != nil {
return nil, trace.Wrap(err)
}
// HTTP server has shorter WriteTimeout than is needed, so we override WriteDeadline of the connection.
if conn, err := authz.ConnFromContext(r.Context()); err == nil {
if err := conn.SetWriteDeadline(h.clock.Now().Add(defaults.HeadlessLoginTimeout)); err != nil {
return nil, trace.Wrap(err)
}
}
loginResp, err := httpClient.AuthenticateSSHUser(r.Context(), authSSHUserReq)
if err != nil {
return nil, trace.Wrap(err)
}
return loginResp, nil
}
// validateTrustedCluster validates the token for a trusted cluster and returns it's own host and user certificate authority.
//
// POST /webapi/trustedclusters/validate
//
// * Request body:
//
// {
// "token": "foo",
// "certificate_authorities": ["AQ==", "Ag=="]
// }
//
// * Response:
//
// {
// "certificate_authorities": ["AQ==", "Ag=="]
// }
func (h *Handler) validateTrustedCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var validateRequestRaw authclient.ValidateTrustedClusterRequestRaw
if err := httplib.ReadJSON(r, &validateRequestRaw); err != nil {
return nil, trace.Wrap(err)
}
validateRequest, err := validateRequestRaw.ToNative()
if err != nil {
return nil, trace.Wrap(err)
}
validateResponse, err := h.auth.ValidateTrustedCluster(r.Context(), validateRequest)
if err != nil {
h.logger.ErrorContext(r.Context(), "Failed validating trusted cluster", "error", err)
if trace.IsAccessDenied(err) {
return nil, trace.AccessDenied("access denied: the cluster token has been rejected")
}
return nil, trace.Wrap(err)
}
validateResponseRaw, err := validateResponse.ToRaw()
if err != nil {
return nil, trace.Wrap(err)
}
return validateResponseRaw, nil
}
func (h *Handler) String() string {
return "multi site"
}
// currentSiteShortcut is a special shortcut that will return the first
// available site, is helpful when UI works in single site mode to reduce
// the amount of requests
const currentSiteShortcut = "-current-"
// ContextHandler is a handler called with the auth context, what means it is authenticated and ready to work
type ContextHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext) (any, error)
// ClusterHandler is a authenticated handler that is called for some existing remote cluster
type ClusterHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error)
// ClusterWebsocketHandler is a authenticated websocket handler that is called for some existing remote cluster
type ClusterWebsocketHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster, ws *websocket.Conn) (any, error)
// WithClusterAuth wraps a ClusterHandler to ensure that a request is authenticated to this proxy
// (the same as WithAuth), as well as to grab the remoteSite (which can represent this local cluster
// or a remote trusted cluster) as specified by the ":site" url parameter.
//
// WithClusterAuth also provides CSRF protection by requiring the bearer token to be present.
func (h *Handler) WithClusterAuth(fn ClusterHandler) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
sctx, site, err := h.authenticateRequestWithCluster(w, r, p)
if err != nil {
return nil, trace.Wrap(err)
}
return fn(w, r, p, sctx, site)
})
}
func (h *Handler) writeErrToWebSocket(ctx context.Context, ws *websocket.Conn, err error) {
if err == nil {
return
}
errEnvelope := terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketError,
Payload: trace.UserMessage(err),
}
env, err := errEnvelope.Marshal()
if err != nil {
h.logger.ErrorContext(ctx, "error marshaling proto", "error", err)
return
}
if err := ws.WriteMessage(websocket.BinaryMessage, env); err != nil {
h.logger.ErrorContext(ctx, "error writing proto", "error", err)
return
}
}
// authnWsUpgrader is an upgrader that allows any origin to connect to the websocket.
// This makes our lives easier in our automated tests. While ordinarily this would be
// used to enforce the same-origin policy, we don't need to worry about that for authenticated
// websockets, which also require a valid bearer token sent over the websocket after upgrade.
// Therefore even if an attacker were to connect to the websocket and trick the browser into
// sending the session cookie, they would still fail to send the bearer token needed to authenticate.
var authnWsUpgrader = websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
CheckOrigin: func(r *http.Request) bool { return true },
// We’re disabling the error handler here since error handling is managed within the handler itself.
// Allowing the WS error handler to operate would result in it writing to the response writer,
// which conflicts with our custom error handler and leads to unintended errors.
Error: func(http.ResponseWriter, *http.Request, int, error) {},
}
// WithClusterAuthWebSocket wraps a ClusterWebsocketHandler to ensure that a request is authenticated
// to this proxy via websocket, as well as to grab the remoteSite (which can represent this local
// cluster or a remote trusted cluster) as specified by the ":site" url parameter.
func (h *Handler) WithClusterAuthWebSocket(fn ClusterWebsocketHandler, opts ...wsUpgraderOption) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
sctx, ws, site, err := h.authenticateWSRequestWithCluster(w, r, p, opts...)
if err != nil {
return nil, trace.Wrap(err)
}
// WS protocol requires the server send a close message
// which should be done by downstream users
defer ws.Close()
if _, err := fn(w, r, p, sctx, site, ws); err != nil {
h.writeErrToWebSocket(r.Context(), ws, err)
}
return nil, nil
})
}
// authenticateWSRequestWithCluster ensures that a request is
// authenticated to this proxy via websocket, returning the
// *SessionContext (same as AuthenticateRequest), and also grabs the
// remoteSite (which can represent this local cluster or a remote
// trusted cluster) as specified by the ":site" url parameter.
func (h *Handler) authenticateWSRequestWithCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params, opts ...wsUpgraderOption) (*SessionContext, *websocket.Conn, reversetunnelclient.Cluster, error) {
sctx, ws, err := h.AuthenticateRequestWS(w, r, opts...)
if err != nil {
return nil, nil, nil, trace.Wrap(err)
}
site, err := h.getSiteByParams(r.Context(), sctx, p)
if err != nil {
return nil, nil, nil, trace.Wrap(err)
}
return sctx, ws, site, nil
}
// authenticateRequestWithCluster ensures that a request is authenticated
// to this proxy, returning the *SessionContext (same as AuthenticateRequest),
// and also grabs the remoteSite (which can represent this local cluster or a
// remote trusted cluster) as specified by the ":site" url parameter.
func (h *Handler) authenticateRequestWithCluster(w http.ResponseWriter, r *http.Request, p httprouter.Params) (*SessionContext, reversetunnelclient.Cluster, error) {
sctx, err := h.AuthenticateRequest(w, r, true)
if err != nil {
return nil, nil, trace.Wrap(err)
}
site, err := h.getSiteByParams(r.Context(), sctx, p)
if err != nil {
return nil, nil, trace.Wrap(err)
}
return sctx, site, nil
}
// getSiteByParams gets the remoteSite (which can represent this local cluster or a
// remote trusted cluster) as specified by the ":site" url parameter.
func (h *Handler) getSiteByParams(ctx context.Context, sctx *SessionContext, p httprouter.Params) (reversetunnelclient.Cluster, error) {
clusterName := p.ByName("site")
site, err := h.getSiteByClusterName(ctx, sctx, clusterName)
if err != nil {
return nil, trace.Wrap(err)
}
return site, nil
}
func (h *Handler) getSiteByClusterName(ctx context.Context, sctx *SessionContext, clusterName string) (reversetunnelclient.Cluster, error) {
if clusterName == currentSiteShortcut {
res, err := h.cfg.ProxyClient.GetClusterName(ctx)
if err != nil {
h.logger.WarnContext(ctx, "Failed to query cluster name", "error", err)
return nil, trace.Wrap(err)
}
clusterName = res.GetClusterName()
}
proxy, err := h.ProxyWithRoles(ctx, sctx)
if err != nil {
h.logger.WarnContext(ctx, "Failed to get proxy with roles", "error", err)
return nil, trace.Wrap(err)
}
cluster, err := proxy.Cluster(ctx, clusterName)
if err != nil {
h.logger.WarnContext(ctx, "Failed to query site", "error", err, "cluster", clusterName)
return nil, trace.Wrap(err)
}
return cluster, nil
}
// ClusterClientProvider is an interface for a type which can provide
// authenticated clients to remote clusters.
type ClusterClientProvider interface {
// UserClientForCluster returns a client to the local or remote cluster
// identified by clusterName and is authenticated with the identity of the
// user.
UserClientForCluster(ctx context.Context, clusterName string) (authclient.ClientI, error)
}
type clusterClientProvider struct {
h *Handler
ctx *SessionContext
}
// UserClientForCluster returns a client to the local or remote cluster
// identified by clusterName and is authenticated with the identity of the user.
func (r clusterClientProvider) UserClientForCluster(ctx context.Context, clusterName string) (authclient.ClientI, error) {
site, err := r.h.getSiteByClusterName(ctx, r.ctx, clusterName)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := r.ctx.GetUserClient(ctx, site)
return clt, trace.Wrap(err)
}
// ClusterClientHandler is an authenticated handler which can get a client for
// any remote cluster.
type ClusterClientHandler func(http.ResponseWriter, *http.Request, httprouter.Params, *SessionContext, ClusterClientProvider) (any, error)
// WithClusterClientProvider wraps a ClusterClientHandler to ensure that a
// request is authenticated to this proxy (the same as WithAuth), and passes a
// ClusterClientProvider so that the handler can access remote clusters. Use
// this instead of WithClusterAuth when the remote cluster cannot be encoded in
// the path or multiple clusters may need to be accessed from a single handler.
func (h *Handler) WithClusterClientProvider(fn ClusterClientHandler) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
sctx, err := h.AuthenticateRequest(w, r, true)
if err != nil {
return nil, trace.Wrap(err)
}
g := clusterClientProvider{
h: h,
ctx: sctx,
}
return fn(w, r, p, sctx, g)
})
}
// ProvisionTokenHandler is a authenticated handler that is called for some existing Token
type ProvisionTokenHandler func(w http.ResponseWriter, r *http.Request, p httprouter.Params, cluster reversetunnelclient.Cluster, token types.ProvisionToken) (any, error)
// WithProvisionTokenAuth ensures that request is authenticated with a provision token.
// Provision tokens, when used like this are invalidated as soon as used.
// Doesn't matter if the underlying response was a success or an error.
func (h *Handler) WithProvisionTokenAuth(fn ProvisionTokenHandler) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
ctx := r.Context()
creds, err := roundtrip.ParseAuthHeaders(r)
if err != nil {
return nil, trace.AccessDenied("need auth")
}
token, err := consumeTokenForAPICall(ctx, h.GetProxyClient(), creds.Password)
if err != nil {
return nil, trace.AccessDenied("need auth")
}
cluster, err := h.cfg.Proxy.Cluster(ctx, h.auth.clusterName)
if err != nil {
h.logger.WarnContext(ctx, "Failed to query cluster", "error", err, "cluster", h.auth.clusterName)
return nil, trace.Wrap(err)
}
return fn(w, r, p, cluster, token)
})
}
// consumeTokenForAPICall will fetch a token, check if the requireRole is present and then delete the token
// If any of those calls returns an error, this method also returns an error
//
// If multiple clients reach here at the same time, only one of them will be able to actually make the request.
// This is possible because the latest call - DeleteToken - returns an error if the resource doesn't exist
// This is currently true for all the backends as explained here
// https://github.com/gravitational/teleport/commit/24fcadc375d8359e80790b3ebeaa36bd8dd2822f
func consumeTokenForAPICall(ctx context.Context, proxyClient authclient.ClientI, tokenName string) (types.ProvisionToken, error) {
token, err := proxyClient.GetToken(ctx, tokenName)
if err != nil {
return nil, trace.Wrap(err)
}
if token.GetJoinMethod() != types.JoinMethodToken {
return nil, trace.BadParameter("unexpected join method %q for token %q", token.GetJoinMethod(), token.GetSafeName())
}
if !checkTokenTTL(token) {
return nil, trace.BadParameter("expired token %q", token.GetSafeName())
}
if err := proxyClient.DeleteToken(ctx, token.GetName()); err != nil {
return nil, trace.Wrap(err)
}
return token, nil
}
// checkTokenTTL returns true if the token is still valid.
// This is similar to checkTokenTTL in auth.Server, but does not delete expired tokens.
func checkTokenTTL(tok types.ProvisionToken) bool {
// Always accept tokens without an expiry configured.
if tok.Expiry().IsZero() {
return true
}
now := time.Now().UTC()
return tok.Expiry().After(now)
}
type redirectHandlerFunc func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (redirectURL string)
// IsValidRedirectURL validates redirect URL.
func IsValidRedirectURL(redirectURL string) bool {
u, err := url.ParseRequestURI(redirectURL)
return err == nil && (!u.IsAbs() || (u.Scheme == "http" || u.Scheme == "https"))
}
// WithRedirect is a handler that redirects to the path specified in the returned value.
func (h *Handler) WithRedirect(fn redirectHandlerFunc) httprouter.Handle {
return func(w http.ResponseWriter, r *http.Request, p httprouter.Params) {
app.SetRedirectPageHeaders(w.Header(), "")
redirectURL := fn(w, r, p)
if !IsValidRedirectURL(redirectURL) {
redirectURL = sso.LoginFailedRedirectURL
}
http.Redirect(w, r, redirectURL, http.StatusFound)
}
}
// WithMetaRedirect is a handler that redirects to the path specified
// using HTML rather than HTTP. This is needed for redirects that can
// have a header size larger than 8kb, which some middlewares will drop.
// See https://github.com/gravitational/teleport/issues/7467.
func (h *Handler) WithMetaRedirect(fn redirectHandlerFunc) httprouter.Handle {
return func(w http.ResponseWriter, r *http.Request, p httprouter.Params) {
redirectURL := fn(w, r, p)
if !IsValidRedirectURL(redirectURL) {
redirectURL = sso.LoginFailedRedirectURL
}
err := app.MetaRedirect(w, redirectURL)
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to issue a redirect", "error", err)
}
}
}
// WithAuth ensures that a request is authenticated.
// Authenticated requests require both a session cookie as well as a bearer token.
// WithAuth also provides CSRF protection by requiring the bearer token to be present.
func (h *Handler) WithAuth(fn ContextHandler) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
sctx, err := h.AuthenticateRequest(w, r, true /* check bearer token */)
if err != nil {
return nil, trace.Wrap(err)
}
return fn(w, r, p, sctx)
})
}
// WithSession ensures that the request provides a session cookie.
// It does not check for a bearer token.
//
// WithSession does not provide CSRF protection, so it should only
// be used for non-state-changing requests or when other CSRF mitigations
// are applied.
func (h *Handler) WithSession(fn ContextHandler) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
sctx, err := h.AuthenticateRequest(w, r, false /* check bearer token */)
if err != nil {
return nil, trace.Wrap(err)
}
return fn(w, r, p, sctx)
})
}
// WithUnauthenticatedLimiter adds a conditional IP-based rate limiting that will limit only unauthenticated requests.
// This is a good default to use as both Cluster and User auth are checked here, but `WithLimiter` can be used if
// you're certain that no authenticated requests will be made.
func (h *Handler) WithUnauthenticatedLimiter(fn httplib.HandlerFunc) httprouter.Handle {
return h.unauthenticatedLimiterFunc(fn, h.limiter)
}
// WithUnauthenticatedHighLimiter adds a conditional IP-based rate limiting that will limit only unauthenticated
// requests. This is similar to WithUnauthenticatedLimiter, however this one allows a much higher rate limit.
// This higher rate limit should only be used on endpoints which are only CPU constrained
// (no file or other resources used).
func (h *Handler) WithUnauthenticatedHighLimiter(fn httplib.HandlerFunc) httprouter.Handle {
return h.unauthenticatedLimiterFunc(fn, h.highLimiter)
}
func (h *Handler) unauthenticatedLimiterFunc(fn httplib.HandlerFunc, limiter *limiter.RateLimiter) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
// check if the request remote is rate limited before attempting authentication.
isLimited, err := isRequestRateLimited(r, limiter)
if err != nil {
return nil, trace.Wrap(err)
}
if isLimited {
return nil, trace.LimitExceeded("rate limit exceeded")
}
if _, _, err := h.authenticateRequestWithCluster(w, r, p); err != nil {
// retry with user auth
if _, err = h.AuthenticateRequest(w, r, true /* check token */); err != nil {
// no auth passed, limit request
return withLimiterHandlerFunc(fn, limiter)(w, r, p)
}
}
// auth passed, call directly
return fn(w, r, p)
})
}
func withLimiterHandlerFunc(fn httplib.HandlerFunc, limiter *limiter.RateLimiter) httplib.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
err := rateLimitRequest(r, limiter)
if err != nil {
return nil, trace.Wrap(err)
}
return fn(w, r, p)
}
}
// WithAccessDeniedLimiter adds conditional IP-based rate limiting that only applies
// when the request handler returns an access denied error.
//
// This is different from WithUnauthenticatedLimiter in an important way:
// - WithUnauthenticatedLimiter: Checks authentication BEFORE executing the handler using local
// authentication mechanisms (user auth session cookie). If authentication fails locally, it
// applies rate limiting BEFORE executing the handler.
// - WithAccessDeniedLimiter: Executes the handler first, delegates authentication
// to the handler, then applies rate limiting ONLY if the handler returns an access denied error.
//
// Use this limiter for endpoints where authentication happens downstream (e.g., SCIM
// endpoints where auth is handled by the auth server, not in the proxy).
func (h *Handler) WithAccessDeniedLimiter(fn httplib.HandlerFunc) httprouter.Handle {
return h.conditionalLimiterFunc(fn, h.limiter, trace.IsAccessDenied)
}
func (h *Handler) conditionalLimiterFunc(fn httplib.HandlerFunc, limiter *limiter.RateLimiter, shouldLimit func(error) bool) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
remoteAddr, err := getRemoteAddressFromRequest(r)
if err != nil {
return nil, trace.Wrap(err)
}
if limiter.IsRateLimited(remoteAddr) {
return nil, trace.LimitExceeded("rate limit exceeded")
}
result, err := fn(w, r, p)
if err != nil {
if shouldLimit(err) {
if err := limiter.RegisterRequest(remoteAddr); err != nil {
return nil, trace.Wrap(err)
}
}
}
return result, err
})
}
// WithLimiter adds IP-based rate limiting to fn.
// Limits are applied to all requests, authenticated or not.
func (h *Handler) WithLimiter(fn httplib.HandlerFunc) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
return h.WithLimiterHandlerFunc(fn)(w, r, p)
})
}
// WithHighLimiter adds high rate IP-based rate limiting to fn.
// This should only be used on functions which are CPU constrained, and don't use disk or other services.
// Limits are applied to all requests, authenticated or not.
func (h *Handler) WithHighLimiter(fn httplib.HandlerFunc) httprouter.Handle {
return httplib.MakeHandler(func(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
return h.WithHighLimiterHandlerFunc(fn)(w, r, p)
})
}
// WithLimiterHandlerFunc adds IP-based rate limiting to a HandlerFunc. This
// should be used when you need to nest this inside another HandlerFunc.
func (h *Handler) WithLimiterHandlerFunc(fn httplib.HandlerFunc) httplib.HandlerFunc {
return withLimiterHandlerFunc(fn, h.limiter)
}
// WithHighLimiterHandlerFunc adds IP-based rate limiting to a HandlerFunc. This is similar to WithLimiterHandlerFunc
// but provides a higher rate limit. This should only be used for requests which are only CPU bound (no disk or other
// resources used).
func (h *Handler) WithHighLimiterHandlerFunc(fn httplib.HandlerFunc) httplib.HandlerFunc {
return withLimiterHandlerFunc(fn, h.highLimiter)
}
// isRequestRateLimited checks if the request would be rate-limited without consuming any tokens.
// Returns true if the request is rate-limited, false otherwise.
func isRequestRateLimited(r *http.Request, limiter *limiter.RateLimiter) (bool, error) {
remote, err := getRemoteAddressFromRequest(r)
if err != nil {
return false, trace.Wrap(err)
}
return limiter.IsRateLimited(remote), nil
}
func getRemoteAddressFromRequest(r *http.Request) (string, error) {
remote, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
return "", trace.Wrap(err)
}
return remote, nil
}
func rateLimitRequest(r *http.Request, limiter *limiter.RateLimiter) error {
remote, err := getRemoteAddressFromRequest(r)
if err != nil {
return trace.Wrap(err)
}
return trace.Wrap(limiter.RegisterRequest(remote))
}
func (h *Handler) validateCookie(w http.ResponseWriter, r *http.Request) (*SessionContext, error) {
const missingCookieMsg = "missing session cookie"
cookie, err := r.Cookie(websession.CookieName)
if err != nil || (cookie != nil && cookie.Value == "") {
return nil, trace.AccessDenied("%s", missingCookieMsg)
}
decodedCookie, err := websession.DecodeCookie(cookie.Value)
if err != nil {
return nil, trace.AccessDenied("failed to decode cookie")
}
sctx, err := h.auth.getOrCreateSession(r.Context(), decodedCookie.User, decodedCookie.SID)
if err != nil {
clearSessionCookies((w))
return nil, trace.AccessDenied("need auth")
}
if sctx.cfg.Session.GetUsage() != types.WebSessionUsage_WEB_SESSION_USAGE_UNSPECIFIED {
clearSessionCookies((w))
return nil, trace.AccessDenied("need auth")
}
return sctx, nil
}
// AuthenticateRequest authenticates request using combination of a session cookie
// and bearer token
func (h *Handler) AuthenticateRequest(w http.ResponseWriter, r *http.Request, checkBearerToken bool) (*SessionContext, error) {
sctx, err := h.validateCookie(w, r)
if err != nil {
return nil, trace.Wrap(err)
}
if checkBearerToken {
creds, err := roundtrip.ParseAuthHeaders(r)
if err != nil {
return nil, trace.AccessDenied("need auth")
}
if err := sctx.validateBearerToken(r.Context(), creds.Password); err != nil {
return nil, trace.AccessDenied("bad bearer token")
}
}
if err := parseMFAResponseFromRequest(r); err != nil {
return nil, trace.Wrap(err)
}
return sctx, nil
}
// AuthenticateReqForAccessGraphAPI is a special authentication method for Access Graph API requests.
// It requires a client TLS certificate with the UsageAccessGraphAPIOnly usage and does not require a bearer token.
// It's used by the enterprise plugin to serve AccessGraph API via CLI.
func (h *Handler) AuthenticateReqForAccessGraphAPI(r *http.Request) (*SessionContext, error) {
if r.TLS == nil || len(r.TLS.PeerCertificates) == 0 {
return nil, trace.AccessDenied("client certificate required")
}
cert := r.TLS.PeerCertificates[0]
identity, err := tlsca.FromSubject(cert.Subject, cert.NotAfter)
if err != nil {
return nil, trace.Wrap(err, "failed to parse client certificate")
}
if len(identity.Usage) != 1 || !slices.Contains(identity.Usage, teleport.UsageAccessGraphAPIOnly) {
return nil, trace.AccessDenied("client certificate is not valid for Access Graph API")
}
if identity.Username == "" || identity.WebSessionID == "" {
return nil, trace.AccessDenied("client certificate missing required fields")
}
sctx, err := h.auth.getOrCreateSession(
r.Context(),
identity.Username,
identity.WebSessionID,
)
if err != nil {
return nil, trace.AccessDenied("need auth")
}
if sctx.cfg.Session.GetUsage() != types.WebSessionUsage_WEB_SESSION_USAGE_ACCESS_GRAPH_API {
return nil, trace.AccessDenied("needs auth")
}
return sctx, nil
}
// parseMFAResponse checks for an MFA response in the request header.
// if found, the mfa response is added to the request context, where
// it can be recalled to augment client authentication further down
// the call stack.
func parseMFAResponseFromRequest(r *http.Request) error {
ctx, err := contextWithMFAResponseFromRequestHeader(r.Context(), r.Header)
if err != nil {
return trace.Wrap(err)
}
// Update the request reference with a cloned request with the update ctx.
*r = *r.WithContext(ctx)
return nil
}
// contextWithMFAResponseFromRequestHeader attempts to parse an MFA response
// from the request header. If found, the MFA response is added to the given
// context and returned.
func contextWithMFAResponseFromRequestHeader(ctx context.Context, requestHeader http.Header) (context.Context, error) {
if mfaResponseJSON := requestHeader.Get("Teleport-MFA-Response"); mfaResponseJSON != "" {
mfaResp, err := client.ParseMFAChallengeResponse([]byte(mfaResponseJSON))
if err != nil {
return nil, trace.Wrap(err)
}
return mfa.ContextWithMFAResponse(ctx, mfaResp), nil
}
return ctx, nil
}
type wsBearerToken struct {
Token string `json:"token"`
}
type wsStatus struct {
Type string `json:"type"`
Status string `json:"status"`
Message string `json:"message,omitempty"`
}
// wsIODeadline is used to set a deadline for receiving a message from
// an authenticated websocket so unauthenticated sockets dont get left
// open.
const wsIODeadline = time.Second * 4
type wsUpgraderOption func(*websocket.Upgrader)
// Adds subprotocol(s) to the websocket upgrade response
func WithSubprotocols(s ...string) wsUpgraderOption {
return func(u *websocket.Upgrader) {
u.Subprotocols = s
}
}
// AuthenticateRequest authenticates request using combination of a session cookie
// and bearer token retrieved from a websocket
func (h *Handler) AuthenticateRequestWS(w http.ResponseWriter, r *http.Request, opts ...wsUpgraderOption) (*SessionContext, *websocket.Conn, error) {
sctx, err := h.validateCookie(w, r)
if err != nil {
return nil, nil, trace.Wrap(err)
}
upgrader := authnWsUpgrader
for _, opt := range opts {
opt(&upgrader)
}
ws, err := upgrader.Upgrade(w, r, nil)
if err != nil {
return nil, nil, trace.ConnectionProblem(err, "Error upgrading to websocket: %v", err)
}
if err := ws.SetReadDeadline(time.Now().Add(wsIODeadline)); err != nil {
return nil, nil, trace.ConnectionProblem(err, "Error setting websocket read deadline: %v", err)
}
var t wsBearerToken
if err := ws.ReadJSON(&t); err != nil {
return nil, nil, trace.Wrap(err)
}
if err := sctx.validateBearerToken(r.Context(), t.Token); err != nil {
writeErr := ws.WriteJSON(wsStatus{
Type: "create_session_response",
Status: "error",
Message: "invalid token",
})
if writeErr != nil {
h.logger.ErrorContext(r.Context(), "Error while writing invalid token error to websocket", "error", writeErr)
}
return nil, nil, trace.Wrap(err)
}
if err := ws.WriteJSON(wsStatus{
Type: "create_session_response",
Status: "ok",
}); err != nil {
return nil, nil, trace.Wrap(err)
}
// unset the deadline as downstream consumers should handle this themselves.
if err := ws.SetReadDeadline(time.Time{}); err != nil {
return nil, nil, trace.ConnectionProblem(err, "Error setting websocket read deadline: %v", err)
}
if err := parseMFAResponseFromRequest(r); err != nil {
return nil, nil, trace.Wrap(err)
}
return sctx, ws, nil
}
// ProxyWithRoles returns a reverse tunnel proxy verifying the permissions
// of the given user.
func (h *Handler) ProxyWithRoles(ctx context.Context, sctx *SessionContext) (reversetunnelclient.ClusterGetter, error) {
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
h.logger.WarnContext(ctx, "Failed to get client roles", "error", err)
return nil, trace.Wrap(err)
}
cn, err := h.cfg.AccessPoint.GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
return reversetunnelclient.NewClusterGetterWithRoles(h.cfg.Proxy, cn.GetClusterName(), accessChecker.CheckAccessToRemoteCluster, h.cfg.AccessPoint), nil
}
// ProxyHostPort returns the address of the proxy server using --proxy
// notation, i.e. "localhost:8030,8023"
func (h *Handler) ProxyHostPort() string {
hp := createHostPort(h.cfg.ProxyWebAddr, defaults.HTTPListenPort)
return fmt.Sprintf("%s,%s", hp, h.sshPort)
}
func (h *Handler) kubeProxyHostPort() string {
return createHostPort(h.cfg.ProxyKubeAddr, defaults.KubeListenPort)
}
// Address can be set in the config with unspecified host, like
// 0.0.0.0:3080 or :443.
//
// In this case, the dial will succeed (dialing 0.0.0.0 is same a dialing
// localhost in Go) but the host certificate will not validate since
// 0.0.0.0 is never a valid principal (auth server explicitly removes it
// when issuing host certs).
//
// As such, replace 0.0.0.0 with localhost in this case: proxy listens on
// all interfaces and localhost is always included in the valid principal
// set.
func createHostPort(netAddr utils.NetAddr, port int) string {
if netAddr.IsHostUnspecified() {
return fmt.Sprintf("localhost:%v", netAddr.Port(port))
}
return netAddr.String()
}
func message(msg string) any {
return map[string]any{"message": msg}
}
// OK is a response that indicates request was successful.
func OK() any {
return message("ok")
}
// makeTeleportClientConfig creates default teleport client configuration
// that is used to initiate an SSH terminal session or SCP file transfer
func makeTeleportClientConfig(ctx context.Context, sctx *SessionContext) (*client.Config, error) {
agent, cert, err := sctx.GetAgent()
if err != nil {
return nil, trace.BadParameter("failed to get user credentials: %v", err)
}
signers, err := agent.Signers()
if err != nil {
return nil, trace.BadParameter("failed to get user credentials: %v", err)
}
tlsConfig, err := sctx.ClientTLSConfig(ctx)
if err != nil {
return nil, trace.BadParameter("failed to get client TLS config: %v", err)
}
callback, err := apisshutils.NewHostKeyCallback(
apisshutils.HostKeyCallbackConfig{
GetHostCheckers: sctx.getCheckers,
Clock: sctx.cfg.Parent.clock,
})
if err != nil {
return nil, trace.Wrap(err)
}
proxyListenerMode, err := sctx.GetProxyListenerMode(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
config := &client.Config{
Username: sctx.GetUser(),
Agent: agent,
NonInteractive: true,
TLS: tlsConfig,
PublicKeyAuthConfig: apissh.PublicKeyAuthConfig{
Signers: func() ([]ssh.Signer, error) {
return signers, nil
},
},
ProxySSHPrincipal: cert.ValidPrincipals[0],
HostKeyCallback: callback,
TLSRoutingEnabled: proxyListenerMode == types.ProxyListenerMode_Multiplex,
Tracer: apitracing.DefaultProvider().Tracer("webterminal"),
AddKeysToAgent: client.AddKeysToAgentNo,
}
return config, nil
}
// SSORequestParams holds parameters parsed out of a HTTP request initiating an
// SSO login. See ParseSSORequestParams().
type SSORequestParams struct {
// ClientRedirectURL is the URL specified in the query parameter
// redirect_url, which will be unescaped here.
ClientRedirectURL string
// ConnectorID identifies the SSO connector to use to log in, from
// the connector_id query parameter.
ConnectorID string
// CSRFToken is used to protect against login-CSRF in SSO flows.
CSRFToken string
// LoginHint is the user's identifier (email) for identifier-first login.
LoginHint string
// Scope is the scope that this session will be pinned to.
Scope string
}
// ParseSSORequestParams extracts the SSO request parameters from an http.Request,
// returning them in an SSORequestParams struct. If any fields are not present,
// an error is returned.
func ParseSSORequestParams(r *http.Request) (*SSORequestParams, error) {
// Manually grab the value from query param "redirect_url".
//
// The "redirect_url" param can contain its own query params such as in
// "https://localhost/login?connector_id=github&redirect_url=https://localhost:8080/web/cluster/im-a-cluster-name/nodes?search=tunnel&sort=hostname:asc",
// which would be incorrectly parsed with the standard method.
// For example a call to r.URL.Query().Get("redirect_url") in the example above would return
// "https://localhost:8080/web/cluster/im-a-cluster-name/nodes?search=tunnel",
// as it would take the "&sort=hostname:asc" to be a separate query param.
//
// This logic assumes that anything coming after "redirect_url" is part of
// the redirect URL.
splittedRawQuery := strings.Split(r.URL.RawQuery, "&redirect_url=")
var clientRedirectURL string
if len(splittedRawQuery) > 1 {
clientRedirectURL, _ = url.QueryUnescape(splittedRawQuery[1])
}
if clientRedirectURL == "" {
return nil, trace.BadParameter("missing redirect_url query parameter")
}
query := r.URL.Query()
loginHint := query.Get("login_hint")
if len(loginHint) > teleport.MaxUsernameLength {
loginHint = ""
}
connectorID := query.Get("connector_id")
if connectorID == "" {
return nil, trace.BadParameter("missing connector_id query parameter")
}
csrfToken, err := csrf.ExtractTokenFromCookie(r)
if err != nil {
return nil, trace.Wrap(err)
}
scope := query.Get("scope")
return &SSORequestParams{
ClientRedirectURL: clientRedirectURL,
ConnectorID: connectorID,
CSRFToken: csrfToken,
LoginHint: loginHint,
Scope: scope,
}, nil
}
// SSOCallbackResponse holds the parameters for validating and executing an SSO
// callback URL. See SSOSetWebSessionAndRedirectURL().
type SSOCallbackResponse struct {
// CSRFToken is the token provided in the originating SSO login request
// to be validated against.
CSRFToken string
// Username is the authenticated teleport username of the user that has
// logged in, provided by the SSO provider.
Username string
// SessionName is the name of the session generated by auth server if
// requested in the SSO request.
SessionName string
// SessionExpiry is the expiration of the session. This is used
// to set the expiration time of the cookie. If no expiraton is set,
// the cookie with be a "session cookie", which is removed when the browser closes.
SessionExpiry time.Time
// ClientRedirectURL is the URL to redirect back to on completion of
// the SSO login process.
ClientRedirectURL string
// MFAToken is an SSO MFA token.
MFAToken string
}
// SSOSetWebSessionAndRedirectURL validates the CSRF token in the response
// against that in the request, validates that the callback URL in the response
// can be parsed, and sets a session cookie with the username and session name
// from the response. On success, nil is returned. If the validation fails, an
// error is returned.
func SSOSetWebSessionAndRedirectURL(w http.ResponseWriter, r *http.Request, response *SSOCallbackResponse, verifyCSRF bool) error {
if verifyCSRF {
// Make sure that the CSRF token provided in this request matches
// the token that was used to initiate this SSO attempt.
//
// This ensures that an attacker cannot perform a login CSRF attack
// in order to get the victim to log in to the incorrect account.
if err := csrf.VerifyToken(response.CSRFToken, r); err != nil {
return trace.Wrap(err)
}
}
if err := websession.SetCookie(w, response.Username, response.SessionName, response.SessionExpiry); err != nil {
return trace.Wrap(err)
}
parsedRedirectURL, err := httplib.OriginLocalRedirectURI(response.ClientRedirectURL)
if err != nil {
return trace.Wrap(err)
}
response.ClientRedirectURL = parsedRedirectURL
return nil
}
const robots = `User-agent: *
Disallow: /`
func serveRobotsTxt(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
w.Header().Set("Content-Type", "text/plain")
w.Header().Set("Cache-Control", "public, max-age=86400")
w.WriteHeader(http.StatusOK)
w.Write([]byte(robots))
return nil, nil
}
func readEtagFromAppHash(fs http.FileSystem) (string, error) {
hashFile, err := fs.Open("/apphash")
if err != nil {
return "", trace.Wrap(err)
}
defer hashFile.Close()
appHash, err := io.ReadAll(hashFile)
if err != nil {
return "", trace.Wrap(err)
}
versionWithHash := fmt.Sprintf("%s-%s", teleport.Version, string(appHash))
etag := fmt.Sprintf("%q", versionWithHash)
return etag, nil
}
/*
* 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 implements web proxy handler that provides
// web interface to view and connect to teleport nodes
package web
import (
"context"
"errors"
"fmt"
"math/rand/v2"
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/web/app"
)
type GetAppDetailsRequest ResolveAppParams
type GetAppDetailsResponse struct {
// FQDN is application FQDN.
FQDN string `json:"fqdn"`
// RequiredAppFQDNs is a list of required app fqdn
RequiredAppFQDNs []string `json:"requiredAppFQDNs"`
}
// getAppDetails resolves the input params to a known application and returns
// its app details.
//
// GET /v1/webapi/apps/:fqdnHint/:clusterName/:publicAddr
func (h *Handler) getAppDetails(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext) (any, error) {
values := r.URL.Query()
isRedirectFlow := values.Get("required-apps") != ""
clusterName := p.ByName("clusterName")
req := GetAppDetailsRequest{
FQDNHint: p.ByName("fqdnHint"),
ClusterName: clusterName,
PublicAddr: p.ByName("publicAddr"),
}
// Use the information the caller provided to attempt to resolve to an
// application running within either the root or leaf cluster.
result, err := h.resolveApp(r.Context(), ctx, ResolveAppParams(req))
if err != nil {
return nil, trace.Wrap(err, "unable to resolve FQDN: %v", req.FQDNHint)
}
resp := &GetAppDetailsResponse{
FQDN: result.FQDN,
}
requiredAppNames := result.App.GetRequiredAppNames()
if !isRedirectFlow {
// TODO (avatus) this would be nice if the string in the RequiredApps spec was the fqdn of the required app
// so we could skip the resolution step all together but this would break existing configs.
// if clusterName is not supplied in the params, the initial app must have been fetched with fqdn hint only.
// We can use the clusterName of the initially resolved app
if clusterName == "" {
clusterName = result.ClusterName
}
// TODO (williamo/scopes): Scoped apps currently won't support required_apps.
if scope := result.App.GetScope(); scope != "" && len(requiredAppNames) > 0 {
return nil, trace.AccessDenied("scoped apps do not support required app redirects")
}
for _, requiredAppName := range requiredAppNames {
if result.App.GetUseAnyProxyPublicAddr() {
proxyDNSName := utils.FindMatchingProxyDNS(req.FQDNHint, h.proxyDNSNames())
requiredAppFQDN := fmt.Sprintf("%s.%s", requiredAppName, proxyDNSName)
resp.RequiredAppFQDNs = append(resp.RequiredAppFQDNs, requiredAppFQDN)
continue
}
res, err := h.resolveApp(r.Context(), ctx, ResolveAppParams{ClusterName: clusterName, AppName: requiredAppName})
if err != nil {
h.logger.ErrorContext(r.Context(), "Error getting app details for associated required app", "required_app", requiredAppName, "app", result.App.GetName(), "error", err)
continue
}
resp.RequiredAppFQDNs = append(resp.RequiredAppFQDNs, res.FQDN)
}
// append self to end of required apps so that it can be the final entry in the redirect "chain".
resp.RequiredAppFQDNs = append(resp.RequiredAppFQDNs, result.FQDN)
}
return resp, nil
}
// CreateAppSessionResponse is a request to POST /v1/webapi/sessions/app
type CreateAppSessionRequest struct {
// ResolveAppParams contains info used to resolve an application
ResolveAppParams
// AWSRole is the AWS role ARN when accessing AWS management console.
AWSRole string `json:"arn,omitempty"`
// MFAResponse is an optional MFA response used to create an MFA verified app session.
MFAResponse client.MFAChallengeResponse `json:"mfaResponse"`
}
// CreateAppSessionResponse is a response to POST /v1/webapi/sessions/app
type CreateAppSessionResponse struct {
// CookieValue is the application session cookie value.
CookieValue string `json:"cookie_value"`
// SubjectCookieValue is the application session subject cookie token.
SubjectCookieValue string `json:"subject_cookie_value"`
// FQDN is application FQDN.
FQDN string `json:"fqdn"`
}
// createAppSession creates a new application session.
//
// POST /v1/webapi/sessions/app
func (h *Handler) createAppSession(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext) (any, error) {
var req CreateAppSessionRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
result, err := h.resolveApp(r.Context(), ctx, req.ResolveAppParams)
if err != nil {
return nil, trace.Wrap(err, "unable to resolve FQDN: %v", req.FQDNHint)
}
h.logger.DebugContext(r.Context(), "Creating application web session", "app_public_addr", result.App.GetPublicAddr(), "cluster", result.ClusterName)
// Ensuring proxy can handle the connection is only done when the request is
// coming from the WebUI.
if h.healthCheckAppServer != nil && !app.HasClientCert(r) {
h.logger.DebugContext(r.Context(), "Ensuring proxy can handle requests for application", "app", result.App.GetName())
err := h.healthCheckAppServer(r.Context(), result.App.GetName(), result.App.GetPublicAddr(), result.ClusterName)
if err != nil {
return nil, trace.ConnectionProblem(err, "Unable to serve application requests. Please try again. If the issue persists, verify if the Application Services are connected to Teleport.")
}
}
mfaResponse, err := req.MFAResponse.GetOptionalMFAResponseProtoReq()
if err != nil {
return nil, trace.Wrap(err)
}
// Get an auth client connected with the user's identity.
authClient, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
// Create an application web session.
//
// Application sessions should not last longer than the parent session.TTL
// will be derived from the identity which has the same expiration as the
// parent web session.
//
// PublicAddr and ClusterName will get encoded within the certificate and
// used for request routing.
ws, err := authClient.CreateAppSession(r.Context(), &proto.CreateAppSessionRequest{
Username: ctx.GetUser(),
PublicAddr: result.App.GetPublicAddr(),
ClusterName: result.ClusterName,
AWSRoleARN: req.AWSRole,
MFAResponse: mfaResponse,
AppName: result.App.GetName(),
URI: result.App.GetURI(),
ClientAddr: r.RemoteAddr,
Scope: result.App.GetScope(),
})
if err != nil {
return nil, trace.Wrap(err)
}
return &CreateAppSessionResponse{
CookieValue: ws.GetName(),
SubjectCookieValue: ws.GetBearerToken(),
FQDN: result.FQDN,
}, nil
}
type ResolveAppParams struct {
// FQDNHint indicates (tentatively) the fully qualified domain name of the application.
FQDNHint string `json:"fqdn,omitempty"`
// PublicAddr is the public address of the application.
PublicAddr string `json:"public_addr,omitempty"`
// ClusterName is the cluster within which this application is running.
ClusterName string `json:"cluster_name,omitempty"`
// AppName is the name of the application
AppName string `json:"app_name,omitempty"`
}
type resolveAppResult struct {
// ServerID is the ID of the server this application is running on.
ServerID string
// FQDN is the best effort FQDN resolved for this application.
FQDN string
// ClusterName is the name of the cluster within which the application
// is running.
ClusterName string
// App is the requested application.
App types.Application
}
// Use the information the caller provided to attempt to resolve to an
// application running within either the root or leaf cluster.
func (h *Handler) resolveApp(ctx context.Context, scx *SessionContext, params ResolveAppParams) (*resolveAppResult, error) {
// Get a reverse tunnel proxy aware of the user's permissions.
proxy, err := h.ProxyWithRoles(ctx, scx)
if err != nil {
return nil, trace.Wrap(err)
}
var (
server types.AppServer
appClusterName string
)
// If the request contains a public address and cluster name (for example, if it came
// from the application launcher in the Web UI) then directly exactly resolve the
// application that the caller is requesting. If it does not, do best effort FQDN resolution.
switch {
case params.ClusterName != "" && (params.AppName != "" || params.PublicAddr != ""):
server, appClusterName, err = h.resolveAppForCluster(ctx, proxy, params.AppName, params.PublicAddr, params.ClusterName)
case params.FQDNHint != "":
// Multiple apps can have the same FQDN - prefer those the user
// can actually access.
var canAccess func(types.Application) bool
if canAccess, err = h.userAppAccessFilter(ctx, scx); err == nil {
server, appClusterName, err = app.ResolveFQDN(ctx, proxy, h.auth.clusterName, h.proxyDNSNames(), params.FQDNHint, canAccess)
}
default:
err = trace.BadParameter("no inputs to resolve application")
}
if err != nil {
return nil, trace.Wrap(err)
}
proxyDNSName := h.proxyDNSName()
if server.GetApp().GetUseAnyProxyPublicAddr() {
proxyDNSName = utils.FindMatchingProxyDNS(params.FQDNHint, h.proxyDNSNames())
}
fqdn := utils.AssembleAppFQDN(h.auth.clusterName, proxyDNSName, appClusterName, server.GetApp())
return &resolveAppResult{
ServerID: server.GetName(),
FQDN: fqdn,
ClusterName: appClusterName,
App: server.GetApp(),
}, nil
}
// userAppAccessFilter returns a filter that checks whether the current user session
// is allowed to access an application.
//
// Per-session MFA and device trust requirements are treated as "accessible" here
// since the user is allowed, they'll just be asked for an extra factor when the connection
// is actually established.
func (h *Handler) userAppAccessFilter(ctx context.Context, scx *SessionContext) (func(types.Application) bool, error) {
checker, err := scx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
authPref, err := h.cfg.AccessPoint.GetAuthPreference(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
state := checker.GetAccessState(authPref)
return func(app types.Application) bool {
err := checker.CheckAccess(app, state)
return err == nil ||
errors.Is(err, services.ErrSessionMFARequired) ||
errors.Is(err, services.ErrTrustedDeviceRequired)
}, nil
}
// resolveAppForCluster will take a cluster name, public address, and optional app name in order to
// locate the application and the server on which it is running.
// The app name, if provided, is used to disambiguate multiple apps with the same public addr.
func (h *Handler) resolveAppForCluster(
ctx context.Context,
clusterGetter reversetunnelclient.ClusterGetter,
appName, publicAddr, clusterName string,
) (types.AppServer, string, error) {
clusterClient, err := clusterGetter.Cluster(ctx, clusterName)
if err != nil {
return nil, "", trace.Wrap(err)
}
servers, err := app.MatchUnshuffled(ctx, clusterClient, app.MatchAppServerForRoute(appName, publicAddr))
if err != nil {
return nil, "", trace.Wrap(err)
}
if len(servers) == 0 {
return nil, "", trace.NotFound("failed to match applications with addr %s and name %q", publicAddr, appName)
}
return servers[rand.N(len(servers))], clusterName, nil
}
// proxyDNSName is a DNS name the HTTP proxy is available at, where
// the local cluster name is used as a best-effort fallback.
func (h *Handler) proxyDNSName() string {
dnsNames := h.proxyDNSNames()
if len(dnsNames) == 0 {
return h.auth.clusterName
}
return dnsNames[0]
}
// proxyDNSNames returns DNS names the HTTP proxy is available at, the local
// cluster name is used as a best-effort fallback.
func (h *Handler) proxyDNSNames() (dnsNames []string) {
for _, addr := range h.cfg.ProxyPublicAddrs {
dnsName, err := utils.DNSName(addr.String())
if err != nil {
continue
}
dnsNames = append(dnsNames, dnsName)
}
if len(dnsNames) == 0 {
return []string{h.auth.clusterName}
}
return dnsNames
}
/*
* Teleport
* Copyright (C) 2024 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"
"fmt"
"net/http"
"strings"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/automaticupgrades/constants"
"github.com/gravitational/teleport/lib/automaticupgrades/version"
)
const defaultChannelTimeout = 5 * time.Second
// automaticUpgrades109 implements a version server in the Teleport Proxy following the RFD 109 spec.
// It is configured through the Teleport Proxy configuration and tells agent updaters
// which version they should install.
func (h *Handler) automaticUpgrades109(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
if h.cfg.AutomaticUpgradesChannels == nil {
return nil, trace.AccessDenied("This proxy is not configured to serve automatic upgrades channels.")
}
// The request format is "<channel name>/{version,critical}"
// As <channel name> might contain "/" we have to split, pop the last part
// and re-construct the channel name.
channelAndType := p.ByName("request")
reqParts := strings.Split(strings.Trim(channelAndType, "/"), "/")
if len(reqParts) < 2 {
return nil, trace.BadParameter("path format should be /webapi/automaticupgrades/channel/<channel>/{version,critical}")
}
requestType := reqParts[len(reqParts)-1]
channelName := strings.Join(reqParts[:len(reqParts)-1], "/")
if channelName == "" {
return nil, trace.BadParameter("a channel name is required")
}
// Finally, we treat the request based on its type
switch requestType {
case "version":
h.logger.DebugContext(r.Context(), "Agent requesting version for channel", "channel", channelName)
return h.automaticUpgradesVersion109(w, r, channelName)
case "critical":
h.logger.DebugContext(r.Context(), "Agent requesting criticality for channel", "channel", channelName)
return h.automaticUpgradesCritical109(w, r, channelName)
default:
return nil, trace.BadParameter("requestType path must end with 'version' or 'critical'")
}
}
// automaticUpgradesVersion109 handles version requests from upgraders
func (h *Handler) automaticUpgradesVersion109(w http.ResponseWriter, r *http.Request, channelName string) (any, error) {
ctx, cancel := context.WithTimeout(r.Context(), defaultChannelTimeout)
defer cancel()
targetVersion, err := h.autoUpdateResolver.GetVersion(ctx, channelName, "" /* updater UUID */)
if err != nil {
// If the error is that the upstream channel has no version
// We gracefully handle by serving "none"
var NoNewVersionErr *version.NoNewVersionError
if errors.As(trace.Unwrap(err), &NoNewVersionErr) {
_, err = w.Write([]byte(constants.NoVersion))
return nil, trace.Wrap(err)
}
// Else we propagate the error
return nil, trace.Wrap(err)
}
// RFD 109 specifies that version from channels must have the leading "v".
// As h.autoUpdateAgentVersion doesn't, we must add it.
_, err = fmt.Fprintf(w, "v%s", targetVersion.String())
return nil, trace.Wrap(err)
}
// automaticUpgradesCritical109 handles criticality requests from upgraders
func (h *Handler) automaticUpgradesCritical109(w http.ResponseWriter, r *http.Request, channelName string) (any, error) {
ctx, cancel := context.WithTimeout(r.Context(), defaultChannelTimeout)
defer cancel()
// RFD109 agents already retrieve maintenance windows from the CMC, no need to
// do a maintenance window lookup for them.
critical, err := h.autoUpdateResolver.ShouldUpdate(ctx, channelName, "" /* updater UUID */, false /* window lookup */)
if err != nil {
return nil, trace.Wrap(err)
}
response := "no"
if critical {
response = "yes"
}
_, err = w.Write([]byte(response))
return nil, trace.Wrap(err)
}
/*
* Teleport
* Copyright (C) 2024 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"
"github.com/coreos/go-semver/semver"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client/webclient"
autoupdatepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/autoupdate/v1"
"github.com/gravitational/teleport/api/types/autoupdate"
"github.com/gravitational/teleport/lib/automaticupgrades/version"
)
// automaticUpdateSettings184 crafts the automatic updates part of the ping/find response
// as described in RFD-184 (agents) and RFD-144 (tools).
func (h *Handler) automaticUpdateSettings184(ctx context.Context, group, updaterUUID string) webclient.AutoUpdateSettings {
// Tools auto updates section.
autoUpdateConfig, err := h.cfg.AccessPoint.GetAutoUpdateConfig(ctx)
// TODO(vapopov) DELETE IN v18.0.0 check of IsNotImplemented, must be backported to all latest supported versions.
if err != nil && !trace.IsNotFound(err) && !trace.IsNotImplemented(err) {
h.logger.ErrorContext(ctx, "failed to receive AutoUpdateConfig", "error", err)
}
autoUpdateVersion, err := h.cfg.AccessPoint.GetAutoUpdateVersion(ctx)
// TODO(vapopov) DELETE IN v18.0.0 check of IsNotImplemented, must be backported to all latest supported versions.
if err != nil && !trace.IsNotFound(err) && !trace.IsNotImplemented(err) {
h.logger.ErrorContext(ctx, "failed to receive AutoUpdateVersion", "error", err)
}
// Agent auto updates section.
agentVersion, err := h.autoUpdateResolver.GetVersion(ctx, group, updaterUUID)
if err != nil {
h.logger.ErrorContext(ctx, "failed to resolve AgentVersion", "error", err)
// Defaulting to current version
agentVersion = teleport.SemVer()
}
// If the source of truth is RFD 109 configuration (channels + CMC) we must emulate the
// RFD109 agent maintenance window behavior by looking up the CMC and checking if
// we are in a maintenance window.
shouldUpdate, err := h.autoUpdateResolver.ShouldUpdate(ctx, group, updaterUUID, true /* window lookup */)
if err != nil {
h.logger.ErrorContext(ctx, "failed to resolve AgentAutoUpdate", "error", err)
// Failing open
shouldUpdate = false
}
toolsVersion, err := getToolsVersion(autoUpdateVersion)
if err != nil {
h.logger.ErrorContext(ctx, "failed to get tools version", "error", err)
toolsVersion = teleport.SemVer()
}
return webclient.AutoUpdateSettings{
ToolsAutoUpdate: getToolsAutoUpdate(autoUpdateConfig),
ToolsVersion: toolsVersion.String(),
AgentUpdateJitterSeconds: DefaultAgentUpdateJitterSeconds,
AgentVersion: agentVersion.String(),
AgentAutoUpdate: shouldUpdate,
}
}
func getToolsAutoUpdate(config *autoupdatepb.AutoUpdateConfig) bool {
// If we can't get the AU config or if AUs are not configured, we default to "disabled".
// This ensures we fail open and don't accidentally update agents if something is going wrong.
// If we want to enable AUs by default, it would be better to create a default "autoupdate_config" resource
// than changing this logic.
if config.GetSpec().GetTools() != nil {
return config.GetSpec().GetTools().GetMode() == autoupdate.ToolsUpdateModeEnabled
}
return false
}
func getToolsVersion(v *autoupdatepb.AutoUpdateVersion) (*semver.Version, error) {
// If we can't get the AU version or tools AU version is not specified, we default to the current proxy version.
// This ensures we always advertise a version compatible with the cluster.
if v.GetSpec().GetTools() == nil {
return teleport.SemVer(), nil
}
return version.EnsureSemver(v.GetSpec().GetTools().GetTargetVersion())
}
/*
* Teleport
* Copyright (C) 2025 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 (
"io"
"mime"
"net/http"
"path"
"slices"
"strings"
)
var compressedFileExtensions = []string{
".js",
".svg",
".wasm",
}
// makeBrotliHandler serves pre-compressed .br files for supported file types.
func makeBrotliHandler(handler http.Handler, fs http.FileSystem) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
ext := path.Ext(r.URL.Path)
isRequestForCompressedFile := slices.Contains(compressedFileExtensions, ext)
clientAcceptsBrotli := strings.Contains(r.Header.Get("Accept-Encoding"), "br")
if !isRequestForCompressedFile || !clientAcceptsBrotli {
handler.ServeHTTP(w, r)
return
}
brPath := r.URL.Path + ".br"
brFile, err := fs.Open(brPath)
if err != nil {
handler.ServeHTTP(w, r)
return
}
defer brFile.Close()
contentType := mime.TypeByExtension(ext)
if contentType == "" {
contentType = "application/octet-stream" // same default as http.DetectContentType
}
w.Header().Set("Content-Encoding", "br")
w.Header().Set("Content-Type", contentType)
if r.Method == http.MethodHead {
return
}
io.Copy(w, brFile)
})
}
// Teleport
// Copyright (C) 2026 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/httplib"
)
// putBrowserMFA accepts a webauthn response from a Browser MFA attempt which is
// sent to CompleteBrowserMFAChallenge for verification. Once verified a tsh
// redirect URL with an encrypted webauthn response is returned.
//
//nolint:staticcheck // TODO(danielashare): Delete when Browser MFA has migrated to mfav2.
func (h *Handler) putBrowserMFA(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
requestID := params.ByName("request_id")
if requestID == "" {
return "", trace.BadParameter("request is missing request ID")
}
var req client.MFAChallengeResponse
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
mfaResp, err := req.GetOptionalMFAResponseProtoReq()
if err != nil {
return nil, trace.Wrap(err)
}
if mfaResp == nil {
return nil, trace.Errorf("mfa response is nil")
}
clt, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := clt.MFAServiceClient().CompleteBrowserMFAChallenge(r.Context(), &mfav1.CompleteBrowserMFAChallengeRequest{
BrowserMfaResponse: &mfav1.BrowserMFAResponse{
RequestId: requestID,
WebauthnResponse: mfaResp.GetWebauthn(),
},
})
if err != nil {
return nil, trace.Wrap(err)
}
return resp.TshRedirectUrl, nil
}
// Teleport
// Copyright (C) 2025 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 (
"archive/zip"
"bytes"
"encoding/json"
"fmt"
"net/http"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/client"
)
// authExportPublic returns the CA Certs that can be used to set up a chain of trust which includes the current Teleport Cluster
//
// GET /webapi/sites/:site/auth/export?type=<auth type>
// GET /webapi/auth/export?type=<auth type>
func (h *Handler) authExportPublic(w http.ResponseWriter, r *http.Request, p httprouter.Params) {
if err := h.authExportPublicError(w, r, p); err != nil {
http.Error(w, err.Error(), trace.ErrorToCode(err))
return
}
// Success output handled by authExportPublicError.
}
// authExportPublicError implements authExportPublic, except it returns an error
// in case of failure. Output is only written on success.
func (h *Handler) authExportPublicError(w http.ResponseWriter, r *http.Request, p httprouter.Params) error {
err := rateLimitRequest(r, h.limiter)
if err != nil {
return trace.Wrap(err)
}
query := r.URL.Query()
caType := query.Get("type") // validated by ExportAllAuthorities
ctx := r.Context()
authorities, err := client.ExportAllAuthorities(
ctx,
h.GetProxyClient(),
client.ExportAuthoritiesRequest{
AuthType: caType,
},
)
if err != nil {
h.logger.DebugContext(ctx, "Failed to generate CA Certs", "error", err)
return trace.Wrap(err)
}
format := query.Get("format")
const formatZip = "zip"
const formatJSON = "json"
switch format {
case "":
break
case formatZip:
return h.authExportPublicZip(w, r, authorities)
case formatJSON:
return h.authExportPublicJSON(w, r, authorities)
default:
return trace.BadParameter("unsupported format %q", format)
}
if l := len(authorities); l > 1 {
return trace.BadParameter("found %d authorities to export, use format=%s or format=%s to export all", l, formatZip, formatJSON)
}
// ServeContent sets the correct headers: Content-Type, Content-Length and Accept-Ranges.
// It also handles the Range negotiation
reader := bytes.NewReader(authorities[0].Data)
http.ServeContent(w, r, "authorized_hosts.txt", time.Now(), reader)
return nil
}
func (h *Handler) authExportPublicZip(
w http.ResponseWriter,
r *http.Request,
authorities []*client.ExportedAuthority,
) error {
now := h.clock.Now().UTC()
// Write authorities to a zip buffer as files named "ca$i.cert".
out := &bytes.Buffer{}
zipWriter := zip.NewWriter(out)
for i, authority := range authorities {
fh := &zip.FileHeader{
Name: fmt.Sprintf("ca%d.cer", i),
Method: zip.Deflate,
Modified: now,
}
fh.SetMode(0644)
fileWriter, err := zipWriter.CreateHeader(fh)
if err != nil {
return trace.Wrap(err)
}
fileWriter.Write(authority.Data)
}
if err := zipWriter.Close(); err != nil {
return trace.Wrap(err)
}
const zipName = "Teleport_CA.zip"
w.Header().Set("Content-Disposition", fmt.Sprintf(`attachment;filename="%s"`, zipName))
http.ServeContent(w, r, zipName, now, bytes.NewReader(out.Bytes()))
return nil
}
func (h *Handler) authExportPublicJSON(
w http.ResponseWriter,
r *http.Request,
authorities []*client.ExportedAuthority,
) error {
marshalledAuthorities, err := json.Marshal(authorities)
if err != nil {
return trace.Wrap(err, "failed to JSON marshal authorities")
}
// File name is not critical here. It is only used by `ServeContent` to determine the value of the
// `Content-Type` header.
http.ServeContent(w, r, "export.json", time.Now(), bytes.NewReader(marshalledAuthorities))
return nil
}
/*
* 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 (
"net/http"
"path"
"slices"
"time"
"github.com/gravitational/teleport/lib/httplib"
)
// makeCacheHandler sets cache headers for cacheable file types.
func makeCacheHandler(handler http.Handler, etag string) http.Handler {
cachedFileTypes := []string{".woff", ".woff2", ".ttf"}
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
// We can cache fonts "permanently" because we don't expect them to change. The rest of our
// assets will have an ETag associated with them (teleport version) that will allow us
// to conditionally send the updated assets or a 304 status (Not Modified) response
if slices.Contains(cachedFileTypes, path.Ext(r.URL.Path)) {
httplib.SetCacheHeaders(w.Header(), time.Hour*24*365 /* one year */)
} else {
httplib.SetEntityTagCacheHeaders(w.Header(), etag)
}
handler.ServeHTTP(w, r)
})
}
/*
* 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"
"log/slog"
"net"
"net/http"
"slices"
"sync"
"time"
"github.com/gobwas/ws"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// connectionUpgrade handles connection upgrades.
func (h *Handler) connectionUpgrade(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
upgrades := r.Header.Values(constants.WebAPIConnUpgradeHeader)
if !slices.Contains(upgrades, constants.WebAPIConnUpgradeTypeWebSocket) {
return nil, trace.NotFound("unsupported upgrade types: %v", upgrades)
}
return h.upgradeALPNWebSocket(w, r, h.upgradeALPN)
}
func (h *Handler) upgradeALPNWebSocket(w http.ResponseWriter, r *http.Request, upgradeHandler ConnectionHandler) (any, error) {
upgrader := websocket.Upgrader{
CheckOrigin: func(r *http.Request) bool { return true },
Subprotocols: []string{
constants.WebAPIConnUpgradeTypeALPN,
constants.WebAPIConnUpgradeTypeALPNPing,
},
}
wsConn, err := upgrader.Upgrade(w, r, nil)
if err != nil {
h.logger.DebugContext(r.Context(), "Failed to upgrade WebSocket.", "error", err)
return nil, trace.Wrap(err)
}
defer wsConn.Close()
h.logger.Log(r.Context(), logutils.TraceLevel, "Received WebSocket upgrade.", "protocol", wsConn.Subprotocol())
// websocketALPNServerConn uses "github.com/gobwas/ws" on the raw net.Conn
// instead of gorilla's websocket.Conn to workaround an issue that
// websocket.Conn caches read error when websocketALPNServerConn is passed
// to a HTTP server and get hijacked for another upgrade. Note that client
// side's (api/client) websocket ALPN connection wrapper also uses
// "github.com/gobwas/ws".
conn := newWebSocketALPNServerConn(r.Context(), wsConn.NetConn(), h.logger)
ctx, cancel := context.WithCancel(r.Context())
defer cancel()
switch wsConn.Subprotocol() {
case constants.WebAPIConnUpgradeTypeALPNPing:
// Starts native WebSocket ping for "alpn-ping".
go h.startPing(ctx, conn)
case constants.WebAPIConnUpgradeTypeALPN:
// Nothing to do
default:
// Just close the connection. Upgrader hijacks the connection so no
// point returning an error.
h.logger.DebugContext(ctx, "Unknown or empty WebSocket subprotocol.", "client_protocols", websocket.Subprotocols(r))
return nil, nil
}
if err := upgradeHandler(ctx, conn); err != nil && !utils.IsOKNetworkError(err) {
// Upgrader hijacks the connection so no point returning an error here.
h.logger.ErrorContext(ctx, "Failed to handle WebSocket upgrade request",
"protocol", wsConn.Subprotocol(),
"error", err,
"remote_addr", logutils.StringerAttr(conn.RemoteAddr()),
)
}
return nil, nil
}
func (h *Handler) upgradeALPN(ctx context.Context, conn net.Conn) error {
if h.cfg.ALPNHandler == nil {
return trace.BadParameter("missing ALPNHandler")
}
// ALPNHandler may handle some connections asynchronously. Here we want to
// block until the handling is done by waiting until the connection is
// closed.
waitConn := newWaitConn(ctx, conn)
defer waitConn.WaitForClose()
return h.cfg.ALPNHandler(ctx, waitConn)
}
type pingWriter interface {
WritePing() error
}
func (h *Handler) startPing(ctx context.Context, pingConn pingWriter) {
ticker := time.NewTicker(defaults.ProxyPingInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
err := pingConn.WritePing()
if err != nil {
if !utils.IsOKNetworkError(err) {
h.logger.WarnContext(ctx, "Failed to write ping message.", "error", err)
}
return
}
}
}
}
// waitConn is a net.Conn that provides a "WaitForClose" function to wait until
// the connection is closed.
type waitConn struct {
net.Conn
ctx context.Context
cancel context.CancelFunc
}
// newWaitConn creates a new waitConn.
func newWaitConn(ctx context.Context, conn net.Conn) *waitConn {
ctx, cancel := context.WithCancel(ctx)
return &waitConn{
Conn: conn,
ctx: ctx,
cancel: cancel,
}
}
// WaitForClose blocks until the Close() function of this connection is called.
func (conn *waitConn) WaitForClose() {
<-conn.ctx.Done()
}
// Close implements net.Conn.
func (conn *waitConn) Close() error {
err := conn.Conn.Close()
conn.cancel()
return trace.Wrap(err)
}
func (conn *waitConn) NetConn() net.Conn {
return conn.Conn
}
type websocketALPNServerConn struct {
net.Conn
readBuffer []byte
readError error
readMutex sync.Mutex
writeMutex sync.Mutex
logContext context.Context
logger *slog.Logger
}
func newWebSocketALPNServerConn(ctx context.Context, conn net.Conn, logger *slog.Logger) *websocketALPNServerConn {
return &websocketALPNServerConn{
Conn: conn,
logContext: ctx,
logger: logger.With(teleport.ComponentKey, teleport.Component(teleport.ComponentWeb, "alpnws")),
}
}
func (c *websocketALPNServerConn) NetConn() net.Conn {
return c.Conn
}
func (c *websocketALPNServerConn) Read(b []byte) (int, error) {
c.readMutex.Lock()
defer c.readMutex.Unlock()
n, err := c.readLocked(b)
// Timeout errors can be temporary. For example, when this connection is
// passed to the kube TLS server, it may get "hijacked" again. During the
// hijack, the SetReadDeadline is called with a past timepoint to fail this
// Read so that the HTTP server's background read can be stopped. In such
// cases, return the original net.Error and clear the cached read error.
var netError net.Error
if errors.As(err, &netError) && netError.Timeout() {
c.readError = nil
c.logger.Log(c.logContext, logutils.TraceLevel, "Cleared cached read error.", "err", netError)
return n, netError
}
return n, trace.Wrap(err)
}
func (c *websocketALPNServerConn) readLocked(b []byte) (int, error) {
// Stop reading if any previous read err.
if c.readError != nil {
return 0, trace.Wrap(c.readError)
}
if len(c.readBuffer) > 0 {
n := copy(b, c.readBuffer)
if n < len(c.readBuffer) {
c.readBuffer = c.readBuffer[n:]
} else {
c.readBuffer = nil
}
return n, nil
}
for {
frame, err := ws.ReadFrame(c.Conn)
if err != nil {
c.readError = err
return 0, trace.Wrap(err)
}
// All client frames should be masked.
if frame.Header.Masked {
frame = ws.UnmaskFrame(frame)
}
c.logger.Log(c.logContext, logutils.TraceLevel, "Read websocket frame.", "op", frame.Header.OpCode, "payload_len", len(frame.Payload))
switch frame.Header.OpCode {
case ws.OpClose:
return 0, io.EOF
case ws.OpBinary:
c.readBuffer = frame.Payload
return c.readLocked(b)
case ws.OpPong:
// Receives Pong as response to Ping. Nothing to do.
}
}
}
func (c *websocketALPNServerConn) writeFrame(frame ws.Frame) error {
c.logger.Log(c.logContext, logutils.TraceLevel, "Writing websocket frame.", "op", frame.Header.OpCode, "payload_len", len(frame.Payload))
c.writeMutex.Lock()
defer c.writeMutex.Unlock()
// There is no need to mask from server to client.
return trace.Wrap(ws.WriteFrame(c.Conn, frame))
}
func (c *websocketALPNServerConn) Write(b []byte) (n int, err error) {
binaryFrame := ws.NewBinaryFrame(b)
if err := c.writeFrame(binaryFrame); err != nil {
return 0, trace.Wrap(err)
}
return len(b), nil
}
func (c *websocketALPNServerConn) WritePing() error {
pingFrame := ws.NewPingFrame([]byte(teleport.ComponentTeleport))
return trace.Wrap(c.writeFrame(pingFrame))
}
func (c *websocketALPNServerConn) SetDeadline(t time.Time) error {
c.writeMutex.Lock()
defer c.writeMutex.Unlock()
return trace.Wrap(c.Conn.SetDeadline(t))
}
func (c *websocketALPNServerConn) SetWriteDeadline(t time.Time) error {
c.writeMutex.Lock()
defer c.writeMutex.Unlock()
return trace.Wrap(c.Conn.SetWriteDeadline(t))
}
func (c *websocketALPNServerConn) SetReadDeadline(t time.Time) error {
c.writeMutex.Lock()
defer c.writeMutex.Unlock()
return trace.Wrap(c.Conn.SetReadDeadline(t))
}
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/client/conntest"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
// getConnectionDiagnostic returns a connection diagnostic connection diagnostics.
func (h *Handler) getConnectionDiagnostic(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
connectionID := p.ByName("connectionid")
connectionDiagnostic, err := clt.GetConnectionDiagnostic(r.Context(), connectionID)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.ConnectionDiagnostic{
ID: connectionDiagnostic.GetName(),
Success: connectionDiagnostic.IsSuccess(),
Message: connectionDiagnostic.GetMessage(),
Traces: ui.ConnectionDiagnosticTraceUIFromTypes(connectionDiagnostic.GetTraces()),
}, nil
}
// diagnoseConnection executes and returns a connection diagnostic.
func (h *Handler) diagnoseConnection(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
req := conntest.TestConnectionRequest{}
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
userClt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
proxySettings, err := h.cfg.ProxySettings.GetProxySettings(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
connectionTesterConfig := conntest.ConnectionTesterConfig{
ResourceKind: req.ResourceKind,
UserClient: userClt,
ProxyHostPort: h.ProxyHostPort(),
PublicProxyAddr: h.PublicProxyAddr(),
KubernetesPublicProxyAddr: h.kubeProxyHostPort(),
TLSRoutingEnabled: proxySettings.TLSRoutingEnabled,
}
tester, err := conntest.ConnectionTesterForKind(connectionTesterConfig)
if err != nil {
return nil, trace.Wrap(err)
}
connectionDiagnostic, err := tester.TestConnection(r.Context(), req)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.ConnectionDiagnostic{
ID: connectionDiagnostic.GetName(),
Success: connectionDiagnostic.IsSuccess(),
Message: connectionDiagnostic.GetMessage(),
Traces: ui.ConnectionDiagnosticTraceUIFromTypes(connectionDiagnostic.GetTraces()),
}, nil
}
/*
* 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 (
"net/http"
"slices"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/connectmycomputer"
"github.com/gravitational/teleport/lib/web/ui"
)
// connectMyComputerLoginsList is a handler for GET /webapi/connectmycomputer/logins.
func (h *Handler) connectMyComputerLoginsList(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext) (any, error) {
connectMyComputerRoleName := connectmycomputer.GetRoleNameForUser(sctx.GetUser())
identity, err := sctx.GetIdentity()
if err != nil {
return nil, trace.Wrap(err)
}
if !slices.Contains(identity.Groups, connectMyComputerRoleName) {
return nil, trace.NotFound("User %s does not have the %s role in the session cert.", sctx.GetUser(), connectMyComputerRoleName)
}
authClient, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
role, err := authClient.GetRole(r.Context(), connectMyComputerRoleName)
if err != nil {
// The user is always able to read roles that they hold, see lib/auth.ServerWithRoles.GetRole.
// Because of this, we don't have to worry about access denied here.
//
// NotFound is also not a factor here. If the role exists in the cert but it has since been
// removed from the cluster, the auth server will respond with access denied.
return nil, trace.Wrap(err, "fetching %q role", connectMyComputerRoleName)
}
logins := role.GetLogins(types.Allow)
return ui.ConnectMyComputerLoginsListResponse{
Logins: logins,
}, nil
}
/*
* 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"
"crypto/sha1"
"crypto/tls"
"encoding/base64"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"time"
gogoproto "github.com/gogo/protobuf/proto"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
oteltrace "go.opentelemetry.io/otel/trace"
apiclient "github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/api/utils/tlsutils"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/client"
dbrepl "github.com/gravitational/teleport/lib/client/db/repl"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/gcp"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/session"
alpncommon "github.com/gravitational/teleport/lib/srv/alpnproxy/common"
dbiam "github.com/gravitational/teleport/lib/srv/db/common/iam"
"github.com/gravitational/teleport/lib/ui"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/web/scripts"
"github.com/gravitational/teleport/lib/web/terminal"
webui "github.com/gravitational/teleport/lib/web/ui"
)
// createOrOverwriteDatabaseRequest contains the necessary basic information
// to create (or overwrite) a database.
// Database here is the database resource, containing information to a real
// database (protocol, uri).
type createOrOverwriteDatabaseRequest struct {
Name string `json:"name,omitempty"`
Labels []ui.Label `json:"labels,omitempty"`
Protocol string `json:"protocol,omitempty"`
URI string `json:"uri,omitempty"`
AWSRDS *awsRDS `json:"awsRds,omitempty"`
// Overwrite will replace an existing db resource
// with a new db resource. Only the name cannot
// be changed.
Overwrite bool `json:"overwrite,omitempty"`
}
type awsRDS struct {
AccountID string `json:"accountId,omitempty"`
ResourceID string `json:"resourceId,omitempty"`
Subnets []string `json:"subnets,omitempty"`
VPCID string `json:"vpcId,omitempty"`
}
func (r *createOrOverwriteDatabaseRequest) checkAndSetDefaults() error {
if r.Name == "" {
return trace.BadParameter("missing database name")
}
if r.Protocol == "" {
return trace.BadParameter("missing protocol")
}
if r.URI == "" {
return trace.BadParameter("missing uri")
}
if r.AWSRDS != nil {
if r.AWSRDS.ResourceID == "" {
return trace.BadParameter("missing aws rds field resource id")
}
if r.AWSRDS.AccountID == "" {
return trace.BadParameter("missing aws rds field account id")
}
if len(r.AWSRDS.Subnets) == 0 {
return trace.BadParameter("missing aws rds field subnets")
}
if r.AWSRDS.VPCID == "" {
return trace.BadParameter("missing aws rds field vpc id")
}
}
return nil
}
// handleDatabaseCreate creates a database's metadata.
func (h *Handler) handleDatabaseCreateOrOverwrite(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var req *createOrOverwriteDatabaseRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
database, err := getNewDatabaseResource(*req)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
if req.Overwrite {
if _, err := clt.GetDatabase(r.Context(), req.Name); err != nil {
return nil, trace.Wrap(err)
}
if err := clt.UpdateDatabase(r.Context(), database); err != nil {
return nil, trace.Wrap(err)
}
} else {
if err := clt.CreateDatabase(r.Context(), database); err != nil {
if trace.IsAlreadyExists(err) {
return nil, trace.AlreadyExists("failed to create database (%q already exists), please use another name", req.Name)
}
return nil, trace.Wrap(err)
}
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
return webui.MakeDatabase(database, accessChecker, h.cfg.DatabaseREPLRegistry, false /* requiresRequest */), nil
}
// updateDatabaseRequest contains some updatable fields of a database resource.
type updateDatabaseRequest struct {
CACert *string `json:"caCert,omitempty"`
Labels []ui.Label `json:"labels,omitempty"`
URI string `json:"uri,omitempty"`
AWSRDS *awsRDS `json:"awsRds,omitempty"`
}
func (r *updateDatabaseRequest) checkAndSetDefaults() error {
if r.CACert != nil {
if *r.CACert == "" {
return trace.BadParameter("missing CA certificate data")
}
if _, err := tlsutils.ParseCertificatePEM([]byte(*r.CACert)); err != nil {
return trace.BadParameter("could not parse provided CA as X.509 PEM certificate")
}
}
// These fields can't be empty if set.
if r.AWSRDS != nil {
if r.AWSRDS.ResourceID == "" {
return trace.BadParameter("missing aws rds field resource id")
}
if r.AWSRDS.AccountID == "" {
return trace.BadParameter("missing aws rds field account id")
}
}
if r.CACert == nil && r.AWSRDS == nil && r.Labels == nil && r.URI == "" {
return trace.BadParameter("missing fields to update the database")
}
return nil
}
// handleDatabaseUpdate updates the database
func (h *Handler) handleDatabasePartialUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
databaseName := p.ByName("database")
if databaseName == "" {
return nil, trace.BadParameter("a database name is required")
}
var req *updateDatabaseRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
database, err := clt.GetDatabase(r.Context(), databaseName)
if err != nil {
return nil, trace.Wrap(err)
}
savedOrNewCaCert := database.GetCA()
if req.CACert != nil {
savedOrNewCaCert = *req.CACert
}
savedOrNewAWSRDS := awsRDS{
AccountID: database.GetAWS().AccountID,
ResourceID: database.GetAWS().RDS.ResourceID,
}
if req.AWSRDS != nil {
savedOrNewAWSRDS = awsRDS{
AccountID: req.AWSRDS.AccountID,
ResourceID: req.AWSRDS.ResourceID,
}
}
savedOrNewURI := req.URI
if len(savedOrNewURI) == 0 {
savedOrNewURI = database.GetURI()
}
savedLabels := database.GetStaticLabels()
// Make a new database to reset the check and set defaulted fields.
database, err = getNewDatabaseResource(createOrOverwriteDatabaseRequest{
Name: databaseName,
Protocol: database.GetProtocol(),
URI: savedOrNewURI,
Labels: req.Labels,
AWSRDS: &savedOrNewAWSRDS,
})
if err != nil {
return nil, trace.Wrap(err)
}
database.SetCA(savedOrNewCaCert)
if len(req.Labels) == 0 {
database.SetStaticLabels(savedLabels)
}
if err := clt.UpdateDatabase(r.Context(), database); err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
return webui.MakeDatabase(database, accessChecker, h.cfg.DatabaseREPLRegistry, false /* requiresRequest */), nil
}
// databaseIAMPolicyResponse is the response type for handleDatabaseGetIAMPolicy.
type databaseIAMPolicyResponse struct {
// Type is the type of the IAM policy.
Type string `json:"type"`
// AWS contains the IAM policy for AWS-hosted databases.
AWS *databaseIAMPolicyAWS `json:"aws,omitempty"`
}
// databaseIAMPolicyAWS contains IAM policy for AWS-hosted databases.
type databaseIAMPolicyAWS struct {
// PolicyDocument is the AWS IAM policy document.
PolicyDocument string `json:"policy_document"`
// Placeholders are placeholders found in the policy document.
Placeholders []string `json:"placeholders,omitempty"`
}
// handleDatabaseGetIAMPolicy returns the required IAM policy for database.
func (h *Handler) handleDatabaseGetIAMPolicy(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
databaseName := p.ByName("database")
if databaseName == "" {
return nil, trace.BadParameter("missing database name")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
dbServers, err := fetchDatabaseServersWithName(r.Context(), clt, r, databaseName)
if err != nil {
return nil, trace.Wrap(err)
}
database := dbServers[0].GetDatabase()
switch {
case database.IsAWSHosted():
policy, placeholders, err := dbiam.GetAWSPolicyDocument(database)
if err != nil {
return nil, trace.Wrap(err)
}
policyJSON, err := json.Marshal(policy)
if err != nil {
return nil, trace.Wrap(err)
}
return &databaseIAMPolicyResponse{
Type: "aws",
AWS: &databaseIAMPolicyAWS{
PolicyDocument: string(policyJSON),
Placeholders: placeholders,
},
}, nil
default:
return nil, trace.BadParameter("IAM policy not supported for database type %q", database.GetType())
}
}
func (h *Handler) sqlServerConfigureADScriptHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
tokenStr := p.ByName("token")
if err := validateJoinToken(tokenStr); err != nil {
return "", trace.Wrap(err)
}
dbAddress := r.URL.Query().Get("uri")
if err := services.ValidateSQLServerURI(dbAddress); err != nil {
return "", trace.BadParameter("invalid database address: %v", err)
}
// verify that the token exists
if _, err := h.GetProxyClient().GetToken(r.Context(), tokenStr); err != nil {
return "", trace.BadParameter("invalid token")
}
proxyServers, err := clientutils.CollectWithFallback(r.Context(), h.GetProxyClient().ListProxyServers, func(context.Context) ([]types.Server, error) {
//nolint:staticcheck // TODO(kiosion) DELETE IN 21.0.0
return h.GetProxyClient().GetProxies()
})
if err != nil {
return "", trace.Wrap(err)
}
if len(proxyServers) == 0 {
return "", trace.NotFound("no proxy servers found")
}
clusterName, err := h.GetProxyClient().GetDomainName(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
certAuthority, err := h.GetProxyClient().GetCertAuthority(
r.Context(),
types.CertAuthID{Type: types.DatabaseClientCA, DomainName: clusterName},
false,
)
if err != nil {
return nil, trace.Wrap(err)
}
caCRL, err := h.GetProxyClient().GenerateCertAuthorityCRL(r.Context(), types.DatabaseClientCA)
if err != nil {
return nil, trace.Wrap(err)
}
if len(certAuthority.GetActiveKeys().TLS) != 1 {
return nil, trace.BadParameter("expected one TLS key pair, got %v", len(certAuthority.GetActiveKeys().TLS))
}
keyPair := certAuthority.GetActiveKeys().TLS[0]
block, _ := pem.Decode(keyPair.Cert)
if block == nil {
return nil, trace.BadParameter("no PEM data in CA data")
}
// Split host and port so we can escape domain characters.
dbHost, dbPort, err := net.SplitHostPort(dbAddress)
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
w.WriteHeader(http.StatusOK)
err = scripts.DatabaseAccessSQLServerConfigureScript.Execute(w, scripts.DatabaseAccessSQLServerConfigureParams{
CACertPEM: string(keyPair.Cert),
CACertSHA1: fmt.Sprintf("%X", sha1.Sum(block.Bytes)),
CACertBase64: base64.StdEncoding.EncodeToString(utils.CreateCertificateBLOB(block.Bytes)),
CRLPEM: string(encodeCRLPEM(caCRL)),
ProxyPublicAddr: proxyServers[0].GetPublicAddr(),
ProvisionToken: tokenStr,
DBAddress: net.JoinHostPort(url.QueryEscape(dbHost), dbPort),
})
return nil, trace.Wrap(err)
}
func (h *Handler) dbConnect(
_ http.ResponseWriter,
r *http.Request,
_ httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (any, error) {
// Create a context for signaling when the terminal session is over and
// link it first with the trace context from the request context
tctx := oteltrace.ContextWithRemoteSpanContext(context.Background(), oteltrace.SpanContextFromContext(r.Context()))
ctx, cancel := context.WithCancel(tctx)
defer cancel()
h.logger.DebugContext(ctx, "Received database interactive connection")
var term session.TerminalParams
q := r.URL.Query()
params := q.Get("params")
if params != "" {
var termReq TerminalRequest
if err := json.Unmarshal([]byte(params), &termReq); err != nil {
h.logger.DebugContext(ctx, "Failed to unmarshal terminal request",
"error", err,
)
} else {
term = termReq.Term
}
}
req, err := readDatabaseSessionRequest(ws)
if err != nil {
if errors.Is(err, io.EOF) || errors.Is(err, net.ErrClosed) || terminal.IsOKWebsocketCloseError(trace.Unwrap(err)) {
h.logger.DebugContext(ctx, "Database interactive session closed before receiving request")
return nil, nil
}
var netError net.Error
if errors.As(trace.Unwrap(err), &netError) && netError.Timeout() {
return nil, trace.BadParameter("timed out waiting for database connect request data on websocket connection")
}
return nil, trace.Wrap(err)
}
log := h.logger.With(
"protocol", req.Protocol,
"service_name", req.ServiceName,
"database_name", req.DatabaseName,
"database_user", req.DatabaseUser,
"database_roles", req.DatabaseRoles,
"remote_addr", logutils.StringerAttr(ws.RemoteAddr()),
)
log.DebugContext(ctx, "Received database interactive session request")
if !h.cfg.DatabaseREPLRegistry.IsSupported(req.Protocol) {
log.ErrorContext(ctx, "Unsupported database protocol")
return nil, trace.NotImplemented("%q database protocol not supported for REPL sessions", req.Protocol)
}
netConfig, err := h.GetAccessPoint().GetClusterNetworkingConfig(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// Get host CA for this Proxy.
clusterName, err := h.GetAccessPoint().GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
proxyHostCA, err := h.GetAccessPoint().GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: clusterName.GetClusterName(),
}, false)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
databases, err := apiclient.GetAllResources[types.DatabaseServer](ctx, clt, &proto.ListResourcesRequest{
Namespace: apidefaults.Namespace,
ResourceType: types.KindDatabaseServer,
PredicateExpression: fmt.Sprintf(`name == %q`, req.ServiceName),
Limit: 1,
})
if err != nil {
return nil, trace.Wrap(err)
}
var db types.Database
if len(databases) > 0 {
db = databases[0].GetDatabase()
adjusted, newUsername := gcp.AdjustDatabaseUsername(req.DatabaseUser, db)
if adjusted {
log.DebugContext(ctx, "Adding default project suffix for IAM principal", "original", req.DatabaseUser, "updated", newUsername)
req.DatabaseUser = newUsername
}
}
sess, err := newDatabaseInteractiveSession(ctx, databaseInteractiveSessionConfig{
log: log,
req: req,
ws: ws,
sctx: sctx,
site: cluster,
clt: clt,
keepAliveInterval: netConfig.GetKeepAliveInterval(),
registry: h.cfg.DatabaseREPLRegistry,
alpnHandler: h.cfg.ALPNHandler,
proxyAddr: h.PublicProxyAddr(),
proxyHostCA: proxyHostCA,
Term: term,
})
if err != nil {
h.logger.ErrorContext(r.Context(), "Failed to create interactive database session", "error", err)
return nil, trace.Wrap(err)
}
defer sess.Close()
if err := sess.Run(); err != nil {
log.ErrorContext(ctx, "Database interactive session exited with error", "error", err)
return nil, trace.Wrap(err)
}
return nil, nil
}
// DatabaseSessionRequest describes a request to create a web-based terminal
// database session.
type DatabaseSessionRequest struct {
// ServiceName is the database resource ID the user will be connected.
ServiceName string `json:"serviceName"`
// Protocol is the database protocol.
Protocol string `json:"protocol"`
// DatabaseName is the database name the session will use.
DatabaseName string `json:"dbName"`
// DatabaseUser is the database user used on the session.
DatabaseUser string `json:"dbUser"`
// DatabaseRoles are ratabase roles that will be attached to the user when
// connecting to the database.
DatabaseRoles []string `json:"dbRoles"`
}
func (r *DatabaseSessionRequest) check() error {
if err := types.ValidateDatabaseName(r.ServiceName); err != nil {
return trace.Wrap(err, "database service name %q is invalid", r.ServiceName)
}
return nil
}
// databaseConnectionRequestWaitTimeout defines how long the server will wait
// for the user to send the connection request.
const databaseConnectionRequestWaitTimeout = defaults.HeadlessLoginTimeout
// readDatabaseSessionRequest reads the database session requestion message from
// websocket connection.
func readDatabaseSessionRequest(ws *websocket.Conn) (*DatabaseSessionRequest, error) {
err := ws.SetReadDeadline(time.Now().Add(databaseConnectionRequestWaitTimeout))
if err != nil {
return nil, trace.Wrap(err, "failed to set read deadline for websocket connection")
}
messageType, bytes, err := ws.ReadMessage()
if err != nil {
return nil, trace.Wrap(err)
}
if err := ws.SetReadDeadline(time.Time{}); err != nil {
return nil, trace.Wrap(err, "failed to set read deadline for websocket connection")
}
if messageType != websocket.BinaryMessage {
return nil, trace.BadParameter("expected binary message of type websocket.BinaryMessage, got %v", messageType)
}
var envelope terminal.Envelope
if err := gogoproto.Unmarshal(bytes, &envelope); err != nil {
return nil, trace.BadParameter("failed to parse envelope: %v", err)
}
if envelope.Type != defaults.WebsocketDatabaseSessionRequest {
return nil, trace.BadParameter("expected database session request but got %q", envelope.Type)
}
var req DatabaseSessionRequest
if err := json.Unmarshal([]byte(envelope.Payload), &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.check(); err != nil {
return nil, trace.Wrap(err)
}
return &req, nil
}
type databaseInteractiveSessionConfig struct {
ws *websocket.Conn
log *slog.Logger
req *DatabaseSessionRequest
sctx *SessionContext
site reversetunnelclient.Cluster
clt authclient.ClientI
keepAliveInterval time.Duration
registry dbrepl.REPLRegistry
alpnHandler ConnectionHandler
proxyAddr string
proxyHostCA types.CertAuthority
// Term is the initial PTY size.
Term session.TerminalParams
}
func (c *databaseInteractiveSessionConfig) check() error {
if c.ws == nil {
return trace.BadParameter("missing parameter ws")
}
if c.req == nil {
return trace.BadParameter("missing parameter req")
}
if c.site == nil {
return trace.BadParameter("missing parameter site")
}
if c.clt == nil {
return trace.BadParameter("missing parameter clt")
}
if c.keepAliveInterval == 0 {
return trace.BadParameter("missing parameter keepAliveInterval")
}
if c.registry == nil {
return trace.BadParameter("missing parameter registry")
}
if c.alpnHandler == nil {
return trace.BadParameter("missing parameter alpnHandler")
}
if c.proxyAddr == "" {
return trace.BadParameter("missing parameter proxyAddr")
}
if c.proxyHostCA == nil {
return trace.BadParameter("missing parameter proxyHostCA")
}
if err := c.Term.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
return nil
}
type databaseInteractiveSession struct {
databaseInteractiveSessionConfig
ctx context.Context
replConn net.Conn
alpnConn net.Conn
stream *terminal.Stream
instance dbrepl.REPLInstance
instanceReadyC chan struct{}
}
func newDatabaseInteractiveSession(ctx context.Context, cfg databaseInteractiveSessionConfig) (*databaseInteractiveSession, error) {
if err := cfg.check(); err != nil {
return nil, trace.Wrap(err)
}
replConn, alpnConn := net.Pipe()
sess := &databaseInteractiveSession{
ctx: ctx,
databaseInteractiveSessionConfig: cfg,
replConn: replConn,
alpnConn: alpnConn,
instanceReadyC: make(chan struct{}),
}
sess.stream = terminal.NewStream(ctx, terminal.StreamConfig{
// Don't close the terminal stream on session error, as it would also
// cause the underlying connection to be closed. This will prevent the
// middleware from properly writing the error into the WebSocket connection.
// The middleware initiates the connection, forwards it to our
// handler, and always closes it.
WS: noopCloserWS{Conn: cfg.ws},
Handlers: map[string]terminal.WSHandlerFunc{
defaults.WebsocketResize: sess.handleWindowResize,
},
})
return sess, nil
}
// noopCloserWS prevents the stream from closing the websocket, to allow the
// middleware to write any returned errors to the client before closing the
// websocket.
type noopCloserWS struct {
*websocket.Conn
}
func (c noopCloserWS) Close() error {
return nil
}
func (s *databaseInteractiveSession) Run() error {
replConn, err := s.makeReplConn()
if err != nil {
return trace.Wrap(err)
}
if err := s.sendSessionMetadata(); err != nil {
return trace.Wrap(err)
}
defaultCloseHandler := s.ws.CloseHandler()
s.ws.SetCloseHandler(func(code int, text string) error {
s.log.DebugContext(s.ctx, "web socket was closed by client - terminating session")
// Call the default close handler if one was set.
if defaultCloseHandler != nil {
err := defaultCloseHandler(code, text)
return trace.NewAggregate(err, s.Close())
}
return trace.Wrap(s.Close())
})
go startWSPingLoop(s.ctx, s.ws, s.keepAliveInterval, s.log, s.Close)
// Wrap s.alpnConn with real client addresses and pass it to the ALPN
// handler.
go func() {
alpnConnWithAddr := utils.NewConnWithAddr(s.alpnConn, s.ws.LocalAddr(), s.ws.RemoteAddr())
if err := s.alpnHandler(s.ctx, alpnConnWithAddr); err != nil && !utils.IsOKNetworkError(err) {
s.log.ErrorContext(s.ctx, "ALPN handler for database interactive session failed", "error", err)
}
}()
repl, err := s.registry.NewInstance(s.ctx, &dbrepl.NewREPLConfig{
Client: s.stream,
ServerConn: replConn,
Route: s.route(),
})
if err != nil {
return trace.Wrap(err)
}
s.instance = repl
if err := repl.SetSize(s.Term.W, s.Term.H); err != nil {
s.log.ErrorContext(s.ctx, "Failed to set initial terminal window size",
"error", err,
)
}
close(s.instanceReadyC)
s.log.DebugContext(s.ctx, "Starting database interactive session")
if err := repl.Run(s.ctx); err != nil {
return trace.Wrap(err)
}
s.log.DebugContext(s.ctx, "Database interactive session exited with success")
return nil
}
func (s *databaseInteractiveSession) Close() error {
// TODO(gabrielcorado): Right now, if we send a close message the UI closes
// the terminal without giving the chance for users to review the session.
// Once this gets solved, we should send the close message here.
if err := s.stream.Close(); err != nil {
s.log.ErrorContext(s.ctx, "Unable to close web socket terminal stream", "error", err)
}
if err := s.replConn.Close(); !utils.IsOKNetworkError(err) {
return trace.Wrap(err)
}
return nil
}
// issueCerts performs the MFA (if required) and generate the user session
// certificates.
func (s *databaseInteractiveSession) issueCerts() (*tls.Certificate, error) {
pk, err := keys.ParsePrivateKey(s.sctx.cfg.Session.GetTLSPriv())
if err != nil {
return nil, trace.Wrap(err, "failed getting user private key from the session")
}
publicKeyPEM, err := keys.MarshalPublicKey(pk.Public())
if err != nil {
return nil, trace.Wrap(err, "failed to marshal public key")
}
routeToDatabase := s.route()
certsReq := proto.UserCertsRequest{
TLSPublicKey: publicKeyPEM,
Username: s.sctx.GetUser(),
Expires: s.sctx.cfg.Session.GetExpiryTime(),
Format: constants.CertificateFormatStandard,
RouteToCluster: s.site.GetName(),
Usage: proto.UserCertsRequest_Database,
RouteToDatabase: routeToDatabase,
}
var certs *proto.Certs
result, err := client.PerformSessionMFACeremony(s.ctx, client.PerformSessionMFACeremonyParams{
CurrentAuthClient: s.clt,
RootAuthClient: s.sctx.cfg.RootClient,
MFACeremony: newMFACeremony(s.stream.WSStream, s.sctx.cfg.RootClient.CreateAuthenticateChallenge, s.proxyAddr),
MFAAgainstRoot: s.sctx.cfg.RootClusterName == s.site.GetName(),
MFARequiredReq: &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_Database{Database: &routeToDatabase},
},
CertsReq: &certsReq,
})
if err != nil && !errors.Is(err, services.ErrSessionMFANotRequired) {
return nil, trace.Wrap(err, "failed performing mfa ceremony")
}
if result != nil {
certs = result.NewCerts
}
if certs == nil {
certs, err = s.sctx.cfg.RootClient.GenerateUserCerts(s.ctx, certsReq)
if err != nil {
return nil, trace.Wrap(err, "failed issuing user certs")
}
}
tlsCert, err := pk.TLSCertificate(certs.TLS)
if err != nil {
return nil, trace.Wrap(err)
}
return &tlsCert, nil
}
// makeReplConn wraps the raw repl conn with a TLS certificate to simulate a
// dialed TLS routing connection.
func (s *databaseInteractiveSession) makeReplConn() (*tls.Conn, error) {
tlsCert, err := s.issueCerts()
if err != nil {
return nil, trace.Wrap(err)
}
alpnProtocol, err := alpncommon.ToALPNProtocol(s.req.Protocol)
if err != nil {
return nil, trace.Wrap(err)
}
proxyAddr, err := utils.ParseAddr(s.proxyAddr)
if err != nil {
return nil, trace.Wrap(err)
}
// The ALPN handler used by the web server was initially intended for ALPN
// connection upgrade. Database handlers serve with the Proxy's host cert on
// the other side.
rootCAs, err := services.CertPool(s.proxyHostCA)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig := &tls.Config{
NextProtos: []string{string(alpnProtocol)},
Certificates: []tls.Certificate{*tlsCert},
RootCAs: rootCAs,
ServerName: proxyAddr.Host(),
}
utils.SetupTLSConfig(tlsConfig, nil /* let server decide cipher */)
return tls.Client(s.replConn, tlsConfig), nil
}
func (s *databaseInteractiveSession) route() proto.RouteToDatabase {
return proto.RouteToDatabase{
Protocol: s.req.Protocol,
ServiceName: s.req.ServiceName,
Username: s.req.DatabaseUser,
Database: s.req.DatabaseName,
Roles: s.req.DatabaseRoles,
}
}
func (s *databaseInteractiveSession) sendSessionMetadata() error {
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: session.Session{
// TODO(gabrielcorado): Have a consistent Session ID. Right now, the
// initial session ID returned won't be correct as the session is only
// initialized by the database server after the REPL starts.
ClusterName: s.site.GetName(),
}})
if err != nil {
return trace.Wrap(err)
}
envelope := &terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketSessionMetadata,
Payload: string(sessionMetadataResponse),
}
envelopeBytes, err := gogoproto.Marshal(envelope)
if err != nil {
return trace.Wrap(err)
}
err = s.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes)
if err != nil {
return trace.Wrap(err)
}
return nil
}
func (s *databaseInteractiveSession) waitForREPLInstance(ctx context.Context) (dbrepl.REPLInstance, error) {
select {
case <-ctx.Done():
return nil, trace.Wrap(ctx.Err())
case <-s.instanceReadyC:
}
if s.instance == nil {
return nil, trace.NotFound("missing database REPL instance")
}
return s.instance, nil
}
func (s *databaseInteractiveSession) handleWindowResize(ctx context.Context, envelope terminal.Envelope) {
repl, err := s.waitForREPLInstance(ctx)
if err != nil {
s.log.DebugContext(ctx, "Failed to get database REPL instance", "error", err)
return
}
if params, err := terminal.ParseWindowResizeMsg(envelope); err != nil {
s.log.WarnContext(ctx, "Failed to handle terminal window resize",
"error", err,
)
} else if err := repl.SetSize(params.W, params.H); err != nil {
s.log.ErrorContext(ctx, "Failed to set terminal window size",
"error", err,
)
}
}
// fetchDatabaseServersWithName fetches all database servers with provided database name.
func fetchDatabaseServersWithName(ctx context.Context, clt resourcesAPIGetter, r *http.Request, databaseName string) ([]types.DatabaseServer, error) {
resp, err := clt.ListResources(ctx, proto.ListResourcesRequest{
Limit: defaults.MaxIterationLimit,
ResourceType: types.KindDatabaseServer,
PredicateExpression: fmt.Sprintf(`name == %q`, databaseName),
UseSearchAsRoles: r.URL.Query().Get("searchAsRoles") == "yes",
})
if err != nil {
return nil, trace.Wrap(err)
}
servers, err := types.ResourcesWithLabels(resp.Resources).AsDatabaseServers()
if err != nil {
return nil, trace.Wrap(err)
}
if len(servers) == 0 {
return nil, trace.NotFound("database %q not found", databaseName)
}
return servers, nil
}
func getNewDatabaseResource(req createOrOverwriteDatabaseRequest) (*types.DatabaseV3, error) {
labels := make(map[string]string)
for _, label := range req.Labels {
labels[label.Name] = label.Value
}
dbSpec := types.DatabaseSpecV3{
Protocol: req.Protocol,
URI: req.URI,
}
if req.AWSRDS != nil {
dbSpec.AWS = types.AWS{
AccountID: req.AWSRDS.AccountID,
RDS: types.RDS{
ResourceID: req.AWSRDS.ResourceID,
Subnets: req.AWSRDS.Subnets,
VPCID: req.AWSRDS.VPCID,
},
}
}
database, err := types.NewDatabaseV3(
types.Metadata{
Name: req.Name,
Labels: labels,
}, dbSpec)
if err != nil {
return nil, trace.Wrap(err)
}
database.SetOrigin(types.OriginDynamic)
return database, nil
}
// encodeCRLPEM takes DER encoded CRL and encodes into PEM.
func encodeCRLPEM(contents []byte) []byte {
return pem.EncodeToMemory(&pem.Block{
Type: "X509 CRL",
Bytes: contents,
})
}
/*
* 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 (
"bytes"
"context"
"crypto"
"crypto/tls"
"errors"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
tdpbv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/desktop/v1"
"github.com/gravitational/teleport/api/mfa"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/client/sso"
"github.com/gravitational/teleport/lib/desktop"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/srv/desktop/tdp"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/legacy"
"github.com/gravitational/teleport/lib/srv/desktop/tdp/protocol/tdpb"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/diagnostics/latency"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
const (
tdpbQueryParameter = "tdpb"
protocolTDP = "teleport-tdp"
)
// GET /webapi/sites/:site/linuxdesktops/:desktopName/connect?username=<username>
func (h *Handler) linuxDesktopConnectHandle(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (any, error) {
desktopName := p.ByName("desktopName")
if desktopName == "" {
return nil, trace.BadParameter("missing desktopName in request URL")
}
log := sctx.cfg.Log.With(
"desktop_name", desktopName,
"cluster_name", cluster.GetName(),
)
log.DebugContext(r.Context(), "New desktop access websocket connection")
if err := h.createDesktopConnection(r, desktopName, log, sctx, cluster, ws, desktop.ConnectToLinuxService, proto.UserCertsRequest_LinuxDesktop); err != nil {
// createDesktopConnection makes a best-effort attempt to send an error to the user
// (via websocket) before terminating the connection. We log the error here, but
// return nil because our HTTP middleware will try to write the returned error in JSON
// format, and this will fail since the HTTP connection has been upgraded to websockets.
log.ErrorContext(r.Context(), "creating desktop connection failed", "error", err)
}
return nil, nil
}
// GET /webapi/sites/:site/desktops/:desktopName/connect?username=<username>
func (h *Handler) desktopConnectHandle(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (any, error) {
desktopName := p.ByName("desktopName")
if desktopName == "" {
return nil, trace.BadParameter("missing desktopName in request URL")
}
log := sctx.cfg.Log.With(
"desktop_name", desktopName,
"cluster_name", cluster.GetName(),
)
log.DebugContext(r.Context(), "New desktop access websocket connection")
if err := h.createDesktopConnection(r, desktopName, log, sctx, cluster, ws, desktop.ConnectToWindowsService, proto.UserCertsRequest_WindowsDesktop); err != nil {
// createDesktopConnection makes a best-effort attempt to send an error to the user
// (via websocket) before terminating the connection. We log the error here, but
// return nil because our HTTP middleware will try to write the returned error in JSON
// format, and this will fail since the HTTP connection has been upgraded to websockets.
log.ErrorContext(r.Context(), "creating desktop connection failed", "error", err)
}
return nil, nil
}
// Adapts a websocket to a tdp.MessageReadWriter.
// Quietly discards TDP messages.
type desktopWebsocketAdapter struct {
conn *websocket.Conn
// Avoid allocating a new byte slice with each received message
// be re-using a buffer.
buf bytes.Buffer
}
// ReadMessage returns a new Message read from the underlying websocket.
func (w *desktopWebsocketAdapter) ReadMessage() (tdp.Message, error) {
for {
w.buf.Reset()
mType, rdr, err := w.conn.NextReader()
if err != nil {
return nil, trace.Wrap(err)
}
if mType != websocket.BinaryMessage {
return nil, trace.Errorf("expected binary message, got: %d", mType)
}
if _, err := io.Copy(&w.buf, rdr); err != nil {
return nil, trace.Wrap(err)
}
msg, err := tdpb.DecodeWithTDPDiscard(w.buf.Bytes())
if err != nil {
if errors.Is(err, tdpb.ErrIsTDP) {
continue
}
return nil, trace.Wrap(err)
}
return msg, nil
}
}
// WriteMessage writes a new Message to the underlying websocket.
func (w *desktopWebsocketAdapter) WriteMessage(m tdp.Message) error {
data, err := m.Encode()
if err != nil {
return trace.Wrap(err)
}
return trace.Wrap(w.conn.WriteMessage(websocket.BinaryMessage, data))
}
// implements handshaker for legacy TDP clients
// TODO(rhammonds) DELETE IN v20.0.0
type tdpHandshaker struct {
connection tdp.MessageReadWriter
withheld []tdp.Message
screenSpec legacy.ClientScreenSpec
// May or may not be nil. Not all web client versions will send a keyboard layout.
keyboardLayout *legacy.ClientKeyboardLayout
}
func (t *tdpHandshaker) sendError(ctx context.Context, log *slog.Logger, err error) error {
if err == nil {
log.WarnContext(ctx, "SendError called with empty message")
err = errors.New("an unknown error has occurred")
}
return trace.Wrap(t.connection.WriteMessage(&legacy.Alert{
Message: err.Error(),
Severity: legacy.SeverityError,
}))
}
func (t *tdpHandshaker) getPromptBuilder(log *slog.Logger) mfaPromptBuilder {
return legacy.NewTDPMFAPrompt(t.connection, &t.withheld, log)
}
func (t *tdpHandshaker) performInitialHandshake(ctx context.Context, log *slog.Logger) error {
msg, err := t.connection.ReadMessage()
if err != nil {
return trace.Wrap(err)
}
screenSpec, ok := msg.(legacy.ClientScreenSpec)
if !ok {
return trace.BadParameter("client sent unexpected message %T", msg)
}
t.screenSpec = screenSpec
width, height := screenSpec.Width, screenSpec.Height
if width > types.MaxRDPScreenWidth || height > types.MaxRDPScreenHeight {
return trace.BadParameter(
"screen size of %d x %d is greater than the maximum allowed by RDP (%d x %d)",
width, height, types.MaxRDPScreenWidth, types.MaxRDPScreenHeight,
)
}
msg, err = t.connection.ReadMessage()
if err != nil {
return trace.Wrap(err)
}
keyboardLayout, gotKeyboardLayout := msg.(legacy.ClientKeyboardLayout)
if !gotKeyboardLayout {
t.withheld = append(t.withheld, msg)
log.InfoContext(ctx, "client did not send keyboard layout", "message_type", logutils.TypeAttr(msg), "width", width, "height", height)
} else {
t.keyboardLayout = &keyboardLayout
}
return nil
}
func (t *tdpHandshaker) forwardTDP(w io.Writer, username string, forwardKeyboardLayout bool) error {
messages := make([]tdp.Message, 0, 3)
messages = append(messages, legacy.ClientUsername{Username: username})
messages = append(messages, t.screenSpec)
if t.keyboardLayout != nil && forwardKeyboardLayout {
// TDPB clients will always send the keyboard layout with the Client Hello.
messages = append(messages, t.keyboardLayout)
}
return sendAll(w, append(messages, t.withheld...))
}
func (t *tdpHandshaker) forwardTDPB(w io.Writer, username string, _ bool) error {
// Convert to Client Hello
hello := &tdpb.ClientHello{
ScreenSpec: tdpbv1.ClientScreenSpec_builder{
Height: t.screenSpec.Height,
Width: t.screenSpec.Width,
}.Build(),
Username: username,
}
if t.keyboardLayout != nil {
hello.KeyboardLayout = t.keyboardLayout.KeyboardLayout
}
withheld, err := translateAll(t.withheld, tdpb.TranslateToModern)
if err != nil {
return trace.Wrap(err)
}
return trace.Wrap(sendAll(w, append([]tdp.Message{hello}, withheld...)))
}
func translateAll(messages []tdp.Message, translate func(tdp.Message) ([]tdp.Message, error)) ([]tdp.Message, error) {
translated := make([]tdp.Message, 0, len(messages))
for _, msg := range messages {
out, err := translate(msg)
if err != nil {
return nil, trace.Wrap(err)
}
if len(out) > 0 {
translated = append(translated, out...)
}
}
return translated, nil
}
// implements handshaker for TDPB clients
type tdpbHandshaker struct {
connection tdp.MessageReadWriter
withheld []tdp.Message
hello *tdpb.ClientHello
}
func (t *tdpbHandshaker) sendError(ctx context.Context, log *slog.Logger, err error) error {
if err == nil {
log.WarnContext(ctx, "sendError called with empty message")
err = errors.New("an unknown error has occurred")
}
return trace.Wrap(t.connection.WriteMessage(&tdpb.Alert{
Message: err.Error(),
Severity: tdpbv1.AlertSeverity_ALERT_SEVERITY_ERROR,
}))
}
func (t *tdpbHandshaker) getPromptBuilder(log *slog.Logger) mfaPromptBuilder {
return mfaPromptBuilder(tdpb.NewTDPBMFAPrompt(t.connection, &t.withheld, log))
}
func (t *tdpbHandshaker) performInitialHandshake(ctx context.Context, log *slog.Logger) error {
upgrade := legacy.TDPUpgrade{}
err := t.connection.WriteMessage(upgrade)
if err != nil {
return trace.Wrap(err)
}
// Now wait patiently for the client to reply with a CLIENT_HELLO TDPB message
// The ReadWriter implementation is expected to discard any legacy TDP messages
// while waiting for the client hello.
msg, err := t.connection.ReadMessage()
if err != nil {
return trace.Wrap(err)
}
var ok bool
t.hello, ok = msg.(*tdpb.ClientHello)
if !ok {
return trace.Errorf("Expected client hello message but got %T", msg)
}
log.InfoContext(ctx, "Received client hello message", "message", t.hello)
return nil
}
func (t *tdpbHandshaker) forwardTDP(w io.Writer, username string, forwardKeyboardLayout bool) error {
messages := make([]tdp.Message, 0, 3)
messages = append(messages, legacy.ClientUsername{Username: username})
screenSpec := legacy.ClientScreenSpec{
Height: t.hello.ScreenSpec.GetHeight(),
Width: t.hello.ScreenSpec.GetWidth(),
}
messages = append(messages, screenSpec)
if forwardKeyboardLayout {
// TDPB clients will always send the keyboard layout with the Client Hello.
messages = append(messages, legacy.ClientKeyboardLayout{KeyboardLayout: t.hello.KeyboardLayout})
}
withheld, err := translateAll(t.withheld, tdpb.TranslateToLegacy)
if err != nil {
return trace.Wrap(err)
}
return sendAll(w, append(messages, withheld...))
}
func (t *tdpbHandshaker) forwardTDPB(w io.Writer, username string, _ bool) error {
t.hello.Username = username
return trace.Wrap(sendAll(w, append([]tdp.Message{t.hello}, t.withheld...)))
}
func sendAll(w io.Writer, messages []tdp.Message) error {
for _, msg := range messages {
if err := tdp.EncodeTo(w, msg); err != nil {
return trace.Wrap(err)
}
}
return nil
}
type handshaker interface {
sendError(context.Context, *slog.Logger, error) error
getPromptBuilder(*slog.Logger) mfaPromptBuilder
performInitialHandshake(context.Context, *slog.Logger) error
forwardTDP(io.Writer, string, bool) error
forwardTDPB(io.Writer, string, bool) error
}
// creates a handshaker instance that interops with either TDP or TDPB clients
func newHandshaker(protocol string, ws *websocket.Conn) handshaker {
if protocol == tdpb.ProtocolName {
return &tdpbHandshaker{
connection: &desktopWebsocketAdapter{conn: ws},
}
}
// Default to TDP
return &tdpHandshaker{
connection: tdp.NewConn(&WebsocketIO{Conn: ws}, legacy.Decode, legacy.WarningConstructor),
}
}
type mfaPromptBuilder func(string) mfa.PromptFunc
type connectorFunc func(ctx context.Context, config *desktop.ConnectionConfig) (conn net.Conn, version string, err error)
func (h *Handler) createDesktopConnection(
r *http.Request,
desktopName string,
log *slog.Logger,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
connectFunc connectorFunc,
certUsage proto.UserCertsRequest_CertUsage,
) error {
defer ws.Close()
ctx := r.Context()
clusterName := cluster.GetName()
// Client may speak TDP or TDPB. We'll know based on the existence of the 'tdpb' query parameter.
// - If the 'tdpb' query parameter is present, then we'll need to send an upgrade message to the client
// and listen for a client_hello message (while discarding any TDP messages received).
// Note: We *always* upgrade the client connection to TDPB if possible.
// - Otherwise fall back to the "legacy" behavior
//
// After either receiving a client_hello or our initial TDP messages, we can dial the agent which
// *also* might speak TDP or TDPB. Unlike the client, the agent only speaks one or the other so we'll
// translate on its behalf if needed.
clientProtocol, err := readClientProtocol(r)
if err != nil {
log.ErrorContext(ctx, "Error reading client desktop protocol", "error", err)
return trace.Wrap(err)
}
log.InfoContext(ctx, "Creating Desktop connection", "client_protocol", clientProtocol)
handshaker := newHandshaker(clientProtocol, ws)
// Read the initial set of TDP messages, or handle TDP upgrade and subsequent
// Client Hello message.
err = handshaker.performInitialHandshake(ctx, log)
if err != nil {
return handshaker.sendError(ctx, log, err)
}
username, err := readUsername(r)
if err != nil {
return handshaker.sendError(ctx, log, err)
}
// Parse the private key of the user from the session context.
pk, err := keys.ParsePrivateKey(sctx.cfg.Session.GetTLSPriv())
if err != nil {
return handshaker.sendError(ctx, log, err)
}
// Check if MFA is required and create a UserCertsRequest.
mfaRequired, certsReq, err := h.prepareForCertIssuance(ctx, sctx, cluster, pk.Public(), desktopName, username, certUsage)
if err != nil {
return handshaker.sendError(ctx, log, err)
}
// Issue certificate for the user/desktop combination and perform MFA ceremony if required.
certs, err := h.issueCerts(ctx, sctx, mfaRequired, certsReq, handshaker.getPromptBuilder(log))
if err != nil {
return handshaker.sendError(ctx, log, err)
}
// Create a TLS config for connecting to the Windows Desktop Service.
tlsConfig, err := h.createDesktopTLSConfig(ctx, sctx, desktopName, pk, certs)
if err != nil {
return handshaker.sendError(ctx, log, err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return handshaker.sendError(ctx, log, err)
}
log.DebugContext(ctx, "Attempting to connect to agent")
clientSrcAddr, clientDstAddr := authz.ClientAddrsFromContext(ctx)
serviceConn, version, err := connectFunc(ctx, &desktop.ConnectionConfig{
Log: log,
DesktopsGetter: clt,
Cluster: cluster,
ClientSrcAddr: clientSrcAddr,
ClientDstAddr: clientDstAddr,
DesktopName: desktopName,
ClusterName: clusterName,
})
if err != nil {
return handshaker.sendError(ctx, log, trace.Wrap(err, "cannot connect to Windows Desktop Service"))
}
defer serviceConn.Close()
serviceConnTLS := tls.Client(serviceConn, tlsConfig)
if err := serviceConnTLS.HandshakeContext(ctx); err != nil {
return handshaker.sendError(ctx, log, err)
}
// ALPN informs us which dialect the server will be using.
// Now that we have a connection to the Windows Desktop Service, we can
// forward the client_hello message (TDPB) or username and screen spec (TDP)
// to the service, and any withheld messages that were received before the MFA
// ceremony was completed.
serverProtocol := serviceConnTLS.ConnectionState().NegotiatedProtocol
switch serverProtocol {
case "":
serverProtocol = protocolTDP
sendKeyboardLayout, _ := utils.MinVerWithoutPreRelease(version, "18.0.0")
err = handshaker.forwardTDP(serviceConnTLS, username, sendKeyboardLayout)
case tdpb.ProtocolName:
err = handshaker.forwardTDPB(serviceConnTLS, username, true /* unused */)
default:
err = trace.BadParameter("Unknown desktop agent protocol %v", serverProtocol)
}
log.InfoContext(ctx, "Connected to agent", "protocol", serverProtocol)
if err != nil {
return handshaker.sendError(ctx, log, err)
}
// this blocks until the connection is closed
handleDesktopWebsocketProxyErr(
ctx,
desktopWebsocketProxy{
ws,
serviceConnTLS,
version,
clientProtocol,
serverProtocol,
log,
}.run(ctx),
log,
)
return nil
}
const (
// SNISuffix is the server name suffix used during SNI to specify the
// target desktop to connect to. The client (proxy_service) will use SNI
// like "${UUID}.desktop.teleport.cluster.local" to pass the UUID of the
// desktop.
// This is a copy of the same constant in `lib/srv/desktop/desktop.go` to
// prevent depending on `lib/srv` in `lib/web`.
SNISuffix = ".desktop." + constants.APIDomain
)
func createUserCertsRequest(sctx *SessionContext, publicKey crypto.PublicKey, desktopName, username, siteName string, certUsage proto.UserCertsRequest_CertUsage) (*proto.UserCertsRequest, error) {
tlsCert, err := sctx.GetX509Certificate()
if err != nil {
return nil, trace.Wrap(err)
}
publicKeyPEM, err := keys.MarshalPublicKey(publicKey)
if err != nil {
return nil, trace.Wrap(err)
}
certsReq := proto.UserCertsRequest{
TLSPublicKey: publicKeyPEM,
Username: tlsCert.Subject.CommonName,
Expires: tlsCert.NotAfter,
RouteToCluster: siteName,
}
certsReq.Usage = certUsage
if certUsage == proto.UserCertsRequest_LinuxDesktop {
certsReq.RouteToLinuxDesktop = proto.RouteToLinuxDesktop{
LinuxDesktop: desktopName,
Login: username,
}
} else {
certsReq.RouteToWindowsDesktop = proto.RouteToWindowsDesktop{
WindowsDesktop: desktopName,
Login: username,
}
}
return &certsReq, nil
}
// prepareForCertIssuance prepares for certificate issuance by checking if MFA
// is required for the user/desktop combination and creating a UserCertsRequest.
func (h *Handler) prepareForCertIssuance(
ctx context.Context,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
publicKey crypto.PublicKey,
desktopName, username string,
certUsage proto.UserCertsRequest_CertUsage,
) (mfaRequired bool, certsReq *proto.UserCertsRequest, err error) {
// Check if MFA is required for this user/desktop combination.
var mfaRequest *IsMFARequiredRequest
if certUsage == proto.UserCertsRequest_LinuxDesktop {
mfaRequest = &IsMFARequiredRequest{
LinuxDesktop: &isMFARequiredLinuxDesktop{
DesktopName: desktopName,
Login: username,
},
}
} else {
mfaRequest = &IsMFARequiredRequest{
WindowsDesktop: &isMFARequiredWindowsDesktop{
DesktopName: desktopName,
Login: username,
},
}
}
mfaRequired, err = h.checkMFARequired(ctx, mfaRequest, sctx, cluster)
if err != nil {
return false, nil, trace.Wrap(err)
}
certsReq, err = createUserCertsRequest(sctx, publicKey, desktopName, username, cluster.GetName(), certUsage)
if err != nil {
return false, nil, trace.Wrap(err)
}
return mfaRequired, certsReq, nil
}
// issueCerts issues certificates for the user/desktop combination, performing
// the MFA ceremony if required.
func (h *Handler) issueCerts(
ctx context.Context,
sctx *SessionContext,
mfaRequired bool,
certsReq *proto.UserCertsRequest,
promptConstructor mfaPromptBuilder,
) (certs *proto.Certs, err error) {
if mfaRequired {
certs, err = h.performSessionMFACeremony(ctx, sctx, certsReq, promptConstructor)
if err != nil {
return nil, trace.Wrap(err)
}
} else {
certs, err = sctx.cfg.RootClient.GenerateUserCerts(ctx, *certsReq)
if err != nil {
return nil, trace.Wrap(err)
}
}
return certs, nil
}
// createDesktopTLSConfig creates a TLS config for connecting to a Windows Desktop Service
// using the user's private key and the issued certificates.
func (h *Handler) createDesktopTLSConfig(
ctx context.Context,
sctx *SessionContext,
desktopName string,
pk *keys.PrivateKey,
certs *proto.Certs,
) (*tls.Config, error) {
certConf, err := pk.TLSCertificate(certs.TLS)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig, err := sctx.ClientTLSConfig(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig.Certificates = []tls.Certificate{certConf}
tlsConfig.NextProtos = []string{tdpb.ProtocolName}
// Pass target desktop name via SNI.
tlsConfig.ServerName = desktopName + SNISuffix
return tlsConfig, nil
}
// performSessionMFACeremony completes the mfa ceremony and returns the raw TLS certificate
// on success. The user will be prompted to tap their security key by the UI
// in order to perform the assertion.
func (h *Handler) performSessionMFACeremony(
ctx context.Context,
sctx *SessionContext,
certsReq *proto.UserCertsRequest,
promptConstructor mfaPromptBuilder,
) (_ *proto.Certs, err error) {
ctx, span := h.tracer.Start(ctx, "desktop/performSessionMFACeremony")
defer func() {
span.RecordError(err)
span.End()
}()
// channelID is used by the front end to differentiate between separate ongoing SSO challenges.
channelID := uuid.NewString()
mfaCeremony := &mfa.Ceremony{
CreateAuthenticateChallenge: sctx.cfg.RootClient.CreateAuthenticateChallenge,
MFACeremonyConstructor: func(_ context.Context) (mfa.CallbackCeremony, error) {
u, err := url.Parse(sso.WebMFARedirect)
if err != nil {
return nil, trace.Wrap(err)
}
u.RawQuery = url.Values{"channel_id": {channelID}}.Encode()
return &sso.MFACeremony{
ClientCallbackURL: u.String(),
ProxyAddress: h.PublicProxyAddr(),
}, nil
},
PromptConstructor: func(po ...mfa.PromptOpt) mfa.Prompt {
return promptConstructor(channelID)
},
}
result, err := client.PerformSessionMFACeremony(ctx, client.PerformSessionMFACeremonyParams{
CurrentAuthClient: nil, // Only RootAuthClient is used.
RootAuthClient: sctx.cfg.RootClient,
MFACeremony: mfaCeremony,
MFAAgainstRoot: true,
MFARequiredReq: nil, // No need to verify.
CertsReq: certsReq,
KeyRing: nil, // We just want the certs.
})
if err != nil {
return nil, trace.Wrap(err)
}
return result.NewCerts, nil
}
func readUsername(r *http.Request) (string, error) {
q := r.URL.Query()
username := q.Get("username")
if username == "" {
return "", trace.BadParameter("missing username in URL")
}
return username, nil
}
func readClientProtocol(r *http.Request) (string, error) {
q := r.URL.Query()
tdpbVersion := q.Get(tdpbQueryParameter)
switch tdpbVersion {
case "":
return protocolTDP, nil
case tdpb.ProtocolName:
return tdpb.ProtocolName, nil
default:
return "", trace.BadParameter("unknown TDPB version %q", tdpbVersion)
}
}
// desktopPinger measures latency between proxy and the desktop by sending legacy.Ping messages
// Windows Desktop Service and measuring the time it takes to receive message with the same UUID back.
type desktopPinger struct {
server tdp.MessageWriter
client tdp.MessageWriter
// when false, the interceptor function swallows ping messages
// without writing to the channel
latencySupported bool
ch chan []byte
}
func (d desktopPinger) intercept(msg tdp.Message) ([]tdp.Message, error) {
var uuid []byte
switch m := msg.(type) {
case legacy.Ping:
uuid = m.UUID[:]
case *tdpb.Ping:
uuid = m.Uuid
default:
// This may be some other legacy TDP message
return []tdp.Message{msg}, nil
}
if !d.latencySupported {
return nil, trace.BadParameter("received unexpected Ping message from server (this is a bug)")
}
d.ch <- uuid
// We've handled the ping. Do not pass it along to the proxy.
return nil, nil
}
func (d desktopPinger) ping(ctx context.Context, ping []byte, msg tdp.Message) error {
// The provided 'ping' byte slice should match the UUID contained in 'msg'
if err := d.server.WriteMessage(msg); err != nil {
return trace.Wrap(err)
}
for {
select {
case pong := <-d.ch:
if bytes.Equal(ping, pong) {
return nil
}
case <-ctx.Done():
return trace.Wrap(ctx.Err())
}
}
}
func (d desktopPinger) reportTDPB(_ context.Context, stats latency.Statistics) error {
return d.client.WriteMessage(&tdpb.LatencyStats{
ClientLatencyMs: uint32(stats.Client),
ServerLatencyMs: uint32(stats.Server),
})
}
func (d desktopPinger) reportTDP(_ context.Context, stats latency.Statistics) error {
return d.client.WriteMessage(legacy.LatencyStats{
ClientLatency: uint32(stats.Client),
ServerLatency: uint32(stats.Server)},
)
}
func (d desktopPinger) pingTDP(ctx context.Context) error {
ping := uuid.New()
return d.ping(ctx, ping[:], legacy.Ping{UUID: ping})
}
func (d desktopPinger) pingTDPB(ctx context.Context) error {
uuid := uuid.New()
return d.ping(ctx, uuid[:], &tdpb.Ping{
Uuid: uuid[:],
})
}
func newConn(rwc io.ReadWriteCloser, protocol string) *tdp.Conn {
if protocol == tdpb.ProtocolName {
return tdp.NewConn(rwc, tdp.DecoderAdapter(tdpb.DecodePermissive), tdpb.WarningConstructor)
}
return tdp.NewConn(rwc, legacy.Decode, legacy.WarningConstructor)
}
type desktopWebsocketProxy struct {
// Client websocket connection
ws *websocket.Conn
// Desktop agent connection
wds net.Conn
// Version of the Desktop Agent
version string
// Client protocol (TDP/TDPB)
clientProtocol string
// Server protocol (TDP/TDPB)
serverProtocol string
log *slog.Logger
}
// run does a bidrectional copy between the websocket
// connection to the browser (ws) and the mTLS connection to Windows
// Desktop Serivce (wds)
func (p desktopWebsocketProxy) run(ctx context.Context) error {
ctx, cancel := context.WithCancel(ctx)
defer func() {
cancel()
p.ws.Close()
p.wds.Close()
}()
var err error
latencySupported := p.serverProtocol == tdpb.ProtocolName
if !latencySupported {
latencySupported, err = utils.MinVerWithoutPreRelease(p.version, "17.5.0")
if err != nil {
return trace.Wrap(err)
}
}
// Create a single pair of legacy.Conn instances. legacy.Conn protects the underlying
// streams with a mutex to allow for concurrent writes.
serverConn := tdp.MessageReadWriteCloser(newConn(p.wds, p.serverProtocol))
clientConn := tdp.MessageReadWriteCloser(newConn(&WebsocketIO{Conn: p.ws}, p.clientProtocol))
pinger := desktopPinger{
// The pinger handles translation internally.
server: serverConn,
client: clientConn,
latencySupported: latencySupported,
ch: make(chan []byte),
}
// The ping interceptor is installed on the server connection
// regardless of whether or not translation is needed
serverConn = tdp.NewReadWriteInterceptor(serverConn, pinger.intercept, nil)
// Translation interceptors will be (optionally) installed in the *write* paths of each connection.
needTranslation := p.clientProtocol != p.serverProtocol
if needTranslation {
// Translation is needed
if p.serverProtocol == tdpb.ProtocolName {
p.log.InfoContext(ctx, "Proxying desktop connection with translation", "server_dialect", tdpb.ProtocolName, "client_dialect", protocolTDP)
// Agent speaks TDPB
// Translate to TDPB when writing to the server. Intercept pings when reading from the server.
serverConn = tdp.NewReadWriteInterceptor(serverConn, nil, tdpb.TranslateToModern)
// Client speaks TDP
// Translate to TDP (legacy) when writing to this connection
clientConn = tdp.NewReadWriteInterceptor(clientConn, nil, tdpb.TranslateToLegacy)
} else {
p.log.InfoContext(ctx, "Proxying desktop connection with translation", "server_dialect", protocolTDP, "client_dialect", tdpb.ProtocolName)
// Agent speaks TDP
// Translate to TDPB when reading from this connection.
serverConn = tdp.NewReadWriteInterceptor(serverConn, nil, tdpb.TranslateToLegacy)
// The client speaks TDPB
// Translate to TDPB (modern) when writing to this connection
clientConn = tdp.NewReadWriteInterceptor(clientConn, nil, tdpb.TranslateToModern)
}
} else {
p.log.InfoContext(ctx, "Proxying desktop connection without translation", "dialect", p.serverProtocol)
}
proxy := tdp.NewConnProxy(clientConn, serverConn)
if latencySupported {
// Default to TDPB
pingerFunc := pinger.pingTDPB
reportFunc := pinger.reportTDPB
// Optionally use TDP versions
if p.serverProtocol == protocolTDP {
pingerFunc = pinger.pingTDP
}
if p.clientProtocol == protocolTDP {
reportFunc = pinger.reportTDP
}
go monitorLatency(
ctx,
clockwork.NewRealClock(),
p.ws,
latency.PingerFunc(pingerFunc),
latency.ReporterFunc(reportFunc),
)
}
// Run joins and returns any read, write, or close errors from each side of the
// connection proxy. We can inspect this singular error chain for any "real"
// network errors (as opposed to errors that are expected from a normal teardown).
err = proxy.Run()
if utils.IsOKNetworkError(err) {
err = nil
}
return trace.Wrap(err)
}
// handleDesktopWebsocketProxyErr handles the error returned by desktopWebsocketProxy by
// unwrapping it and determining whether to log an error.
func handleDesktopWebsocketProxyErr(ctx context.Context, proxyWsConnErr error, log *slog.Logger) {
if proxyWsConnErr == nil {
log.DebugContext(ctx, "desktopWebsocketProxy returned with no error")
return
}
errs := []error{proxyWsConnErr}
for len(errs) > 0 {
err := errs[0] // pop first error
errs = errs[1:]
var aggregateErr trace.Aggregate
var closeErr *websocket.CloseError
switch {
case errors.As(err, &aggregateErr):
errs = append(errs, aggregateErr.Errors()...)
case errors.As(err, &closeErr):
switch closeErr.Code {
case websocket.CloseNormalClosure, // when the user hits "disconnect" from the menu
websocket.CloseGoingAway: // when the user closes the tab
log.DebugContext(ctx, "Web socket closed by client", "close_code", closeErr.Code)
return
}
return
default:
if wrapped := errors.Unwrap(err); wrapped != nil {
errs = append(errs, wrapped)
}
}
}
log.WarnContext(ctx, "Error proxying a desktop protocol websocket to windows_desktop_service", "error", proxyWsConnErr)
}
/*
* 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"
"net/http"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/player"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/web/desktop"
)
func (h *Handler) desktopPlaybackHandle(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (any, error) {
sID := p.ByName("sid")
if sID == "" {
return nil, trace.BadParameter("missing session ID in request URL")
}
sessionID, err := session.ParseID(sID)
if err != nil {
return nil, trace.BadParameter("invalid session ID in request URL - %v", err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
player, err := player.New(&player.Config{
Clock: h.clock,
Log: h.logger,
SessionID: *sessionID,
Streamer: clt,
Context: r.Context(),
})
if err != nil {
h.logger.ErrorContext(r.Context(), "couldn't create player for session", "session_id", sID, "error", err)
ws.WriteMessage(websocket.BinaryMessage,
[]byte(`{"message": "error", "errorText": "Internal server error"}`))
return nil, nil
}
defer player.Close()
ctx, cancel := context.WithCancel(r.Context())
defer cancel()
go func() {
defer cancel()
err := desktop.ReceivePlaybackActions(ctx, h.logger, ws, player)
// Connection close errors are expected if the user closes the tab.
// Only log unexpected errors to avoid cluttering the logs.
if !utils.IsOKNetworkError(err) {
h.logger.WarnContext(ctx, "websocket read error", "error", err)
}
}()
go func() {
defer cancel()
defer ws.Close()
player.Play()
desktop.StreamRecording(ctx, h.logger, ws, player)
}()
<-ctx.Done()
return nil, nil
}
// Teleport
// Copyright (C) 2024 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 (
"net/http"
"net/url"
"path"
"strings"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
devicepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/devicetrust/v1"
"github.com/gravitational/teleport/lib/web/app"
)
// deviceWebConfirm is the last step in device web authentication, where the
// "authenticator process" (aka Connect) forwards the DeviceConfirmationToken
// back to the Auth Server, via the Proxy.
//
// GET /webapi/devices/webconfirm?id=a&token=b
//
// - id: ID of the confirmation token.
// - token: raw confirmation token.
//
// The result of this call is a redirect to "/web", regardless of the outcome of
// the ConfirmDeviceWebAuthentication RPC.
func (h *Handler) deviceWebConfirm(w http.ResponseWriter, r *http.Request, _ httprouter.Params, sessionCtx *SessionContext) (any, error) {
query := r.URL.Query()
// Read input parameters.
confirmToken := &devicepb.DeviceConfirmationToken{}
confirmToken.SetId(query.Get("id"))
confirmToken.SetToken(query.Get("token"))
unsafeRedirectURI := query.Get("redirect_uri")
switch {
case confirmToken.GetId() == "":
return nil, trace.BadParameter("parameter id required")
case confirmToken.GetToken() == "":
return nil, trace.BadParameter("parameter token required")
}
// Use the Proxy identity for this call. Only the Proxy is allowed to do it.
devicesClient := h.GetProxyClient().DevicesClient()
ctx := r.Context()
_, err := devicesClient.ConfirmDeviceWebAuthentication(ctx, devicepb.ConfirmDeviceWebAuthenticationRequest_builder{
ConfirmationToken: confirmToken,
CurrentWebSessionId: sessionCtx.GetSessionID(),
}.Build())
switch {
case err != nil:
h.logger.WarnContext(ctx, "Device web authentication confirm failed",
"error", err,
"user", sessionCtx.GetUser(),
)
// err swallowed on purpose.
default:
// Preemptively release session from cache, as its certificates are now
// updated.
// The WebSession watcher takes care of this in other proxy instances
// (see [sessionCache.watchWebSessions]).
h.auth.releaseResources(r.Context(), sessionCtx.GetUser(), sessionCtx.GetSessionID())
}
// Always redirect back to the dashboard, regardless of outcome.
app.SetRedirectPageHeaders(w.Header(), "" /* nonce */)
redirectTo, err := h.getRedirectURL(r.Host, unsafeRedirectURI)
if err != nil {
h.logger.DebugContext(ctx, "Unable to parse redirectURI",
"error", err,
"redirect_uri", unsafeRedirectURI,
)
http.Error(w, http.StatusText(trace.ErrorToCode(err)), trace.ErrorToCode(err))
return nil, nil
}
http.Redirect(w, r, redirectTo, http.StatusSeeOther)
return nil, nil
}
// getRedirectPath tries to parse the given unsafeRedirectURI.
// It returns a full URL if the unsafeRedirectURI points to SAML IdP SSO endpoint.
// In any other case, as long as the redirect URL is parsable, it returns
// a path ensuring its prefixed with "/web".
//
// Nobody seems to know why we need to prepend the base path to the URL, so we keep doing it. It
// might be related to the URLs we get from SSO redirects [1], but it's unclear why we'd be getting
// a URL that's missing the base path and becomes valid only after appending the base path.
//
// [1]: https://github.com/gravitational/teleport/pull/47221#discussion_r1792248868
func (h *Handler) getRedirectURL(host, unsafeRedirectURI string) (string, error) {
const (
basePath = "/web"
samlSPInitiatedSSOPath = "/enterprise/saml-idp/sso"
samlIDPInitiatedSSOPath = "/enterprise/saml-idp/login"
)
if unsafeRedirectURI == "" {
return basePath, nil
}
parsedURL, err := url.Parse(unsafeRedirectURI)
if err != nil {
return basePath, trace.BadParameter("invalid redirect URL")
}
cleanPath := path.Clean(parsedURL.Path)
// helps in situations where there is no path such as https://example.com
if cleanPath == "." || cleanPath == ".." {
cleanPath = "/"
} else if !strings.HasPrefix(cleanPath, "/") {
cleanPath = "/" + cleanPath
}
// IDP initiated SSO path format: "/enterprise/saml-idp/login/<service provider name>"
isIdpInitiatedSSOPath := strings.HasPrefix(cleanPath, samlIDPInitiatedSSOPath) && len(strings.Split(cleanPath, "/")) == 5
if cleanPath == samlSPInitiatedSSOPath || isIdpInitiatedSSOPath {
if parsedURL.Host != host {
return "", trace.BadParameter("host mismatch")
}
path := samlSPInitiatedSSOPath
if isIdpInitiatedSSOPath {
path = cleanPath
}
ensuredURL := &url.URL{
Scheme: "https",
Host: host,
Path: path,
RawQuery: parsedURL.RawQuery,
}
return ensuredURL.String(), nil
}
// Prepend "/web" only if it's not already present
if !strings.HasPrefix(cleanPath, basePath) {
cleanPath = path.Join(basePath, cleanPath)
}
if parsedURL.RawQuery != "" {
return cleanPath + "?" + parsedURL.RawQuery, nil
}
return cleanPath, nil
}
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types/discoveryconfig"
"github.com/gravitational/teleport/api/types/header"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
// discoveryconfigCreate creates a DiscoveryConfig
func (h *Handler) discoveryconfigCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var req ui.DiscoveryConfig
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
dc, err := discoveryconfig.NewDiscoveryConfig(
header.Metadata{
Name: req.Name,
},
discoveryconfig.Spec{
DiscoveryGroup: req.DiscoveryGroup,
AWS: req.AWS,
Azure: req.Azure,
GCP: req.GCP,
Kube: req.Kube,
AccessGraph: req.AccessGraph,
},
)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
storedDiscoveryConfig, err := clt.DiscoveryConfigClient().CreateDiscoveryConfig(r.Context(), dc)
if err != nil {
if trace.IsAlreadyExists(err) {
return nil, trace.AlreadyExists("failed to create DiscoveryConfig (%q already exists), please use another name", req.Name)
}
return nil, trace.Wrap(err)
}
return ui.MakeDiscoveryConfig(storedDiscoveryConfig), nil
}
// discoveryconfigUpdate updates the DiscoveryConfig based on its name
func (h *Handler) discoveryconfigUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
dcName := p.ByName("name")
if dcName == "" {
return nil, trace.BadParameter("a discoveryconfig name is required")
}
var req *ui.UpdateDiscoveryConfigRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
dc, err := clt.DiscoveryConfigClient().GetDiscoveryConfig(r.Context(), dcName)
if err != nil {
return nil, trace.Wrap(err)
}
dc.Spec.DiscoveryGroup = req.DiscoveryGroup
dc.Spec.AWS = req.AWS
dc.Spec.Azure = req.Azure
dc.Spec.GCP = req.GCP
dc.Spec.Kube = req.Kube
dc.Spec.AccessGraph = req.AccessGraph
dc, err = clt.DiscoveryConfigClient().UpdateDiscoveryConfig(r.Context(), dc)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeDiscoveryConfig(dc), nil
}
// discoveryconfigDelete removes a DiscoveryConfig based on its name
func (h *Handler) discoveryconfigDelete(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
discoveryconfigName := p.ByName("name")
if discoveryconfigName == "" {
return nil, trace.BadParameter("a discoveryconfig name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
if err := clt.DiscoveryConfigClient().DeleteDiscoveryConfig(r.Context(), discoveryconfigName); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// discoveryconfigGet returns a DiscoveryConfig based on its name
func (h *Handler) discoveryconfigGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
discoveryconfigName := p.ByName("name")
if discoveryconfigName == "" {
return nil, trace.BadParameter("as discoveryconfig name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
dc, err := clt.DiscoveryConfigClient().GetDiscoveryConfig(r.Context(), discoveryconfigName)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeDiscoveryConfig(dc), nil
}
// discoveryconfigList returns a page of DiscoveryConfigs
func (h *Handler) discoveryconfigList(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
startKey := values.Get("startKey")
dcs, nextKey, err := clt.DiscoveryConfigClient().ListDiscoveryConfigs(r.Context(), int(limit), startKey)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.DiscoveryConfigsListResponse{
Items: ui.MakeDiscoveryConfigs(dcs),
NextKey: nextKey,
}, nil
}
/*
* Teleport
* Copyright (C) 2024 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"
"github.com/gravitational/teleport/api/client/proto"
)
// SetClusterFeatures sets the flags for supported and unsupported features.
// TODO(mcbattirola): make method unexported, fix tests using it to set
// test modules instead.
func (h *Handler) SetClusterFeatures(features proto.Features) {
h.Mutex.Lock()
defer h.Mutex.Unlock()
h.clusterFeatures = features
}
// GetClusterFeatures returns flags for supported and unsupported features.
func (h *Handler) GetClusterFeatures() proto.Features {
h.Mutex.Lock()
defer h.Mutex.Unlock()
return h.clusterFeatures
}
// startFeatureWatcher periodically pings the auth server and updates `clusterFeatures`.
// Must be called only once per `handler`, otherwise it may close an already closed channel
// which will cause a panic.
// The watcher doesn't ping the auth server immediately upon start because features are
// already set by the config object in `NewHandler`.
func (h *Handler) startFeatureWatcher(ctx context.Context) {
ticker := h.clock.NewTicker(h.cfg.FeatureWatchInterval)
h.logger.InfoContext(ctx, "Proxy handler features watcher has started", "interval", h.cfg.FeatureWatchInterval)
defer ticker.Stop()
for {
select {
case <-ticker.Chan():
h.logger.InfoContext(ctx, "Pinging auth server for features")
pingResponse, err := h.GetProxyClient().Ping(ctx)
if err != nil {
h.logger.ErrorContext(ctx, "Auth server ping failed", "error", err)
continue
}
if pingResponse.ServerFeatures != nil {
h.SetClusterFeatures(*pingResponse.ServerFeatures)
h.logger.InfoContext(ctx, "Done updating proxy features", "features", pingResponse.ServerFeatures)
}
case <-ctx.Done():
h.logger.InfoContext(ctx, "Feature service has stopped")
return
}
}
}
/*
* 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"
"net/http"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/defaults"
tracessh "github.com/gravitational/teleport/api/observability/tracing/ssh"
apissh "github.com/gravitational/teleport/api/ssh"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/api/utils/sshutils"
"github.com/gravitational/teleport/lib/agentless"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/multiplexer"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/sshagent"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/sshutils/sftp"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/session/sftputils"
)
// fileTransferRequest describes HTTP file transfer request
type fileTransferRequest struct {
// Server describes a server to connect to (serverId|hostname[:port]).
serverID string
// Login is Linux username to connect as.
login string
// Cluster is the name of the remote cluster to connect to.
cluster string
// remoteLocation is file remote location
remoteLocation string
// filename is a file name
filename string
// mfaResponse is an optional parameter that contains an mfa response string used to issue single use certs
mfaResponse string
// fileTransferRequestID is used to find a FileTransferRequest on a session
fileTransferRequestID string
// moderatedSessonID is an ID of a moderated session that has completed a
// file transfer request approval process
moderatedSessionID string
}
func (h *Handler) transferFile(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
query := r.URL.Query()
req := fileTransferRequest{
cluster: cluster.GetName(),
login: p.ByName("login"),
serverID: p.ByName("server"),
remoteLocation: query.Get("location"),
filename: query.Get("filename"),
mfaResponse: query.Get("mfaResponse"),
fileTransferRequestID: query.Get("fileTransferRequestId"),
moderatedSessionID: query.Get("moderatedSessionId"),
}
var mfaResponse *proto.MFAAuthenticateResponse
if req.mfaResponse != "" {
var err error
if mfaResponse, err = client.ParseMFAChallengeResponse([]byte(req.mfaResponse)); err != nil {
return nil, trace.Wrap(err)
}
}
// Send an error if only one of these params has been sent. Both should exist or not exist together
if (req.fileTransferRequestID != "") != (req.moderatedSessionID != "") {
return nil, trace.BadParameter("fileTransferRequestId and moderatedSessionId must both be included in the same request.")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ft := fileTransfer{
sctx: sctx,
authClient: clt,
proxyHostPort: h.ProxyHostPort(),
}
mfaReq, err := clt.IsMFARequired(r.Context(), &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_Node{
Node: &proto.NodeLogin{
Node: p.ByName("server"),
Login: p.ByName("login"),
},
},
})
if err != nil {
return nil, trace.Wrap(err)
}
if mfaReq.Required && mfaResponse == nil {
return nil, trace.AccessDenied("MFA required for file transfer")
}
tc, err := ft.createClient(req, r, h.cfg.PROXYSigner)
if err != nil {
return nil, trace.Wrap(err)
}
if req.mfaResponse != "" {
if err = ft.issueSingleUseCert(mfaResponse, r, tc); err != nil {
return nil, trace.Wrap(err)
}
}
var moderatedSessionID string
if req.fileTransferRequestID != "" {
moderatedSessionID = req.moderatedSessionID
}
accessPoint, err := cluster.CachingAccessPoint()
if err != nil {
h.logger.DebugContext(r.Context(), "Unable to get auth access point", "error", err)
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
getAgent := sshagent.NewStaticClientGetter(tc.LocalAgent())
cert, err := sctx.GetSSHCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
ident, err := sshca.DecodeIdentity(cert)
if err != nil {
return nil, trace.Wrap(err)
}
certGen, err := h.cfg.Router.GetSiteClient(ctx, tc.SiteName)
if err != nil {
return nil, trace.Wrap(err)
}
signer := agentless.SignerFromSSHIdentity(ident, h.auth.accessPoint, certGen, tc.SiteName, tc.Username)
clientDstAddr := h.cfg.ProxyWebAddr
if srvConn, err := authz.ConnFromContext(r.Context()); err == nil {
clientDstAddr = utils.FromAddr(srvConn.LocalAddr())
}
conn, err := h.cfg.Router.DialHost(
ctx,
ident.ScopePin,
&utils.NetAddr{Addr: r.RemoteAddr},
&clientDstAddr,
req.serverID,
"0",
tc.SiteName,
accessChecker.CheckAccessToRemoteCluster,
getAgent,
signer,
)
if err != nil {
if errors.Is(err, teleport.ErrNodeIsAmbiguous) {
const message = "error: ambiguous host could match multiple nodes\n\nHint: try addressing the node by unique id (ex: user@node-id)\n"
return nil, trace.NotFound("%s", message)
}
return nil, trace.Wrap(err)
}
dialTimeout := defaults.DefaultIOTimeout
if netConfig, err := accessPoint.GetClusterNetworkingConfig(ctx); err != nil {
h.logger.DebugContext(r.Context(), "Unable to fetch cluster networking config", "error", err)
} else {
dialTimeout = netConfig.GetSSHDialTimeout()
}
sshConfig := apissh.ClientConfig{
User: tc.HostLogin,
PublicKeyAuth: tc.PublicKeyAuthConfig,
HostKeyCallback: tc.HostKeyCallback,
Timeout: dialTimeout,
}
nodeClient, err := client.NewNodeClient(
ctx,
sshConfig,
conn,
req.serverID+":0",
req.serverID,
tc,
h.cfg.Modules.IsFIPSBuild(),
)
if err != nil {
// The close error is ignored instead of using [trace.NewAggregate] because
// aggregate errors do not allow error inspection with things like [trace.IsAccessDenied].
_ = conn.Close()
return nil, trace.Wrap(err)
}
defer nodeClient.Close()
webTarget := sftp.Target{
Path: req.filename,
}
remoteTarget := sftp.Target{
Login: req.login,
Addr: &utils.NetAddr{
Addr: req.serverID + ":0",
},
Path: req.remoteLocation,
}
dialHost := func(_ context.Context, _, _ string) (*tracessh.Client, error) {
return nodeClient.Client, nil
}
var sftpReq *sftp.FileTransferRequest
if r.Method == http.MethodPost {
sftpReq, err = sftp.CreateHTTPUploadRequest(sftp.HTTPTransferRequest{
Src: webTarget,
Dst: remoteTarget,
HTTPRequest: r,
DialHost: dialHost,
ModeratedSessionID: moderatedSessionID,
})
} else {
sftpReq, err = sftp.CreateHTTPDownloadRequest(sftp.HTTPTransferRequest{
Src: remoteTarget,
Dst: webTarget,
HTTPResponse: w,
DialHost: dialHost,
ModeratedSessionID: moderatedSessionID,
})
}
if err != nil {
return nil, trace.Wrap(err)
}
if err := sftp.TransferFiles(ctx, sftpReq); err != nil {
if errors.As(err, new(*sftputils.NonRecursiveDirectoryTransferError)) {
return nil, trace.Errorf("transferring directories through the Web UI is not supported at the moment, please use tsh scp -r")
}
return nil, trace.Wrap(err)
}
// We must return nil so that we don't write anything to
// the response, which would corrupt the downloaded file.
return nil, nil
}
type fileTransfer struct {
// sctx is a web session context for the currently logged in user.
sctx *SessionContext
authClient authclient.ClientI
proxyHostPort string
}
func (f *fileTransfer) createClient(req fileTransferRequest, httpReq *http.Request, proxySigner multiplexer.PROXYHeaderSigner) (*client.TeleportClient, error) {
if req.login == "" {
return nil, trace.BadParameter("missing login")
}
servers, err := f.authClient.GetNodes(httpReq.Context(), defaults.Namespace)
if err != nil {
return nil, trace.Wrap(err)
}
hostName, hostPort, err := resolveServerHostPort(req.serverID, servers)
if err != nil {
return nil, trace.BadParameter("invalid server name %q: %v", req.serverID, err)
}
cfg, err := makeTeleportClientConfig(httpReq.Context(), f.sctx)
if err != nil {
return nil, trace.Wrap(err)
}
cfg.HostLogin = req.login
cfg.SiteName = req.cluster
if err := cfg.ParseProxyHost(f.proxyHostPort); err != nil {
return nil, trace.BadParameter("failed to parse proxy address: %v", err)
}
cfg.Host = hostName
cfg.HostPort = hostPort
cfg.ClientAddr = httpReq.RemoteAddr
cfg.PROXYSigner = proxySigner
tc, err := client.NewClient(cfg)
if err != nil {
return nil, trace.BadParameter("failed to create client: %v", err)
}
return tc, nil
}
// issueSingleUseCert will take an assertion response sent from a solved challenge in the web UI
// and use that to generate a cert. This cert is added to the Teleport Client as an authmethod that
// can be used to connect to a node.
func (f *fileTransfer) issueSingleUseCert(mfaResponse *proto.MFAAuthenticateResponse, httpReq *http.Request, tc *client.TeleportClient) error {
pk, err := keys.ParsePrivateKey(f.sctx.cfg.Session.GetSSHPriv())
if err != nil {
return trace.Wrap(err)
}
// Always acquire certs from the root cluster, that is where both the user and their devices are registered.
cert, err := f.sctx.cfg.RootClient.GenerateUserCerts(httpReq.Context(), proto.UserCertsRequest{
SSHPublicKey: pk.MarshalSSHPublicKey(),
Username: f.sctx.GetUser(),
Expires: time.Now().Add(time.Minute).UTC(),
MFAResponse: mfaResponse,
})
if err != nil {
return trace.Wrap(err)
}
sshCert, err := sshutils.ParseCertificate(cert.SSH)
if err != nil {
return trace.Wrap(err)
}
signer, err := sshutils.SSHSigner(sshCert, pk.Signer)
if err != nil {
return trace.Wrap(err)
}
tc.PublicKeyAuthConfig = apissh.PublicKeyAuthConfig{
Signers: func() ([]ssh.Signer, error) {
return []ssh.Signer{signer}, nil
},
}
return nil
}
/*
* Teleport
* Copyright (C) 2025 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
func (h *Handler) gitServerCreateOrUpsert(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var req *ui.CreateGitServerRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.Check(); err != nil {
return nil, trace.Wrap(err)
}
// Only GitHub server is supported. Above req.Check() performs necessary
// checks to ensure all the fields are set.
gitServer, err := types.NewGitHubServerWithName(req.Name, types.GitHubServerMetadata{
Organization: req.GitHub.Organization,
Integration: req.GitHub.Integration,
})
if err != nil {
return nil, trace.Wrap(err)
}
userClient, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
gitServiceClient := userClient.GitServerClient()
if req.Overwrite {
upserted, err := gitServiceClient.UpsertGitServer(r.Context(), gitServer)
return upserted, trace.Wrap(err)
}
created, err := gitServiceClient.CreateGitServer(r.Context(), gitServer)
return created, trace.Wrap(err)
}
func (h *Handler) gitServerGet(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
name := p.ByName("name")
if name == "" {
return nil, trace.BadParameter("git server name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
gitServer, err := clt.GitServerClient().GetGitServer(r.Context(), name)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeGitServer(cluster.GetName(), gitServer, false), nil
}
func (h *Handler) gitServerDelete(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
name := p.ByName("name")
if name == "" {
return nil, trace.BadParameter("git server name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
if err := clt.GitServerClient().DeleteGitServer(r.Context(), name); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/httplib"
)
const headlessAuthID = "headless_authentication_id"
func (h *Handler) getHeadless(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
headlessAuthenticationID, err := getHeadlessAuthID(params)
if err != nil {
return nil, trace.Wrap(err)
}
authClient, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
headlessAuthn, err := authClient.GetHeadlessAuthentication(r.Context(), headlessAuthenticationID)
if err != nil {
// Log the error, but return something more user-friendly.
// Context exceeded or invalid request states are more confusing than helpful.
h.logger.DebugContext(r.Context(), "failed to get headless session", "error", err)
return nil, trace.BadParameter("requested invalid headless session")
}
return headlessAuthn, nil
}
func (h *Handler) putHeadlessState(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
headlessAuthenticationID, err := getHeadlessAuthID(params)
if err != nil {
return nil, trace.Wrap(err)
}
var req client.HeadlessRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
// MFAResponse is required only when accepting a request.
mfaResp, err := req.MFAResponse.GetOptionalMFAResponseProtoReq()
if err != nil {
return nil, trace.Wrap(err)
}
var action types.HeadlessAuthenticationState
switch req.Action {
case "accept":
action = types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_APPROVED
case "denied":
action = types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_DENIED
default:
return nil, trace.BadParameter("unknown action %s", req.Action)
}
authClient, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
if err := authClient.UpdateHeadlessAuthenticationState(r.Context(), headlessAuthenticationID, action, mfaResp); err != nil {
return nil, trace.Wrap(err)
}
// WebUI expects a JSON response.
return OK(), nil
}
func getHeadlessAuthID(params httprouter.Params) (string, error) {
headlessAuthenticationID := params.ByName(headlessAuthID)
if headlessAuthenticationID == "" {
return "", trace.BadParameter("request is missing headless authentication ID")
}
return headlessAuthenticationID, nil
}
/*
* 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"
"iter"
"log/slog"
"net/http"
"net/url"
"slices"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
discoveryconfigv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/discoveryconfig/v1"
integrationv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/integration/v1"
pluginspb "github.com/gravitational/teleport/api/gen/proto/go/teleport/plugins/v1"
usertasksv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/usertasks/v1"
"github.com/gravitational/teleport/api/mfa"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/discoveryconfig"
"github.com/gravitational/teleport/api/types/usertasks"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/integrations/access/msteams"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
libui "github.com/gravitational/teleport/lib/ui"
"github.com/gravitational/teleport/lib/web/ui"
)
// integrationsCreate creates an Integration
func (h *Handler) integrationsCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var req *ui.CreateIntegrationRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
var ig *types.IntegrationV1
var err error
switch req.SubKind {
case types.IntegrationSubKindAWSOIDC:
var s3Location string
if req.AWSOIDC.IssuerS3Bucket != "" {
issuerS3URI := url.URL{
Scheme: "s3",
Host: req.AWSOIDC.IssuerS3Bucket,
Path: req.AWSOIDC.IssuerS3Prefix,
}
s3Location = issuerS3URI.String()
}
metadata := types.Metadata{Name: req.Name}
ig, err = types.NewIntegrationAWSOIDC(
metadata,
&types.AWSOIDCIntegrationSpecV1{
RoleARN: req.AWSOIDC.RoleARN,
IssuerS3URI: s3Location,
Audience: req.AWSOIDC.Audience,
},
)
if err != nil {
return nil, trace.Wrap(err)
}
case types.IntegrationSubKindGitHub:
ig, err = types.NewIntegrationGitHub(types.Metadata{
Name: req.Name,
}, &types.GitHubIntegrationSpecV1{
Organization: req.Integration.GitHub.Organization,
})
if err != nil {
return nil, trace.Wrap(err)
}
cred := types.PluginCredentialsV1{
Credentials: &types.PluginCredentialsV1_IdSecret{
IdSecret: &types.PluginIdSecretCredential{
Id: req.OAuth.ID,
Secret: req.OAuth.Secret,
},
},
}
if err := ig.SetCredentials(&cred); err != nil {
return nil, trace.Wrap(err)
}
case types.IntegrationSubKindAWSRolesAnywhere:
ig, err = types.NewIntegrationAWSRA(types.Metadata{
Name: req.Name,
}, &types.AWSRAIntegrationSpecV1{
TrustAnchorARN: req.Integration.AWSRA.TrustAnchorARN,
ProfileSyncConfig: &types.AWSRolesAnywhereProfileSyncConfig{
Enabled: req.Integration.AWSRA.ProfileSyncConfig.Enabled,
ProfileARN: req.Integration.AWSRA.ProfileSyncConfig.ProfileARN,
RoleARN: req.Integration.AWSRA.ProfileSyncConfig.RoleARN,
ProfileNameFilters: req.Integration.AWSRA.ProfileSyncConfig.ProfileNameFilters,
},
})
if err != nil {
return nil, trace.Wrap(err)
}
default:
return nil, trace.BadParameter("subkind %q is not supported", req.SubKind)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
storedIntegration, err := clt.CreateIntegration(r.Context(), ig)
if err != nil {
if trace.IsAlreadyExists(err) {
return nil, trace.AlreadyExists("failed to create Integration (%q already exists), please use another name", req.Name)
}
return nil, trace.Wrap(err)
}
uiIg, err := ui.MakeIntegration(storedIntegration)
if err != nil {
return nil, trace.Wrap(err)
}
return uiIg, nil
}
// integrationsUpdate updates the Integration based on its name
func (h *Handler) integrationsUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
var req *ui.UpdateIntegrationRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
// Strip MFA from context for the read call so the MFA response is not
// consumed before the UpdateIntegration call that actually needs it.
getCtx := mfa.ContextWithMFAResponse(r.Context(), nil)
integration, err := clt.GetIntegration(getCtx, integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
if req.AWSOIDC != nil {
if integration.GetSubKind() != types.IntegrationSubKindAWSOIDC {
return nil, trace.BadParameter("cannot update %q fields for a %q integration", types.IntegrationSubKindAWSOIDC, integration.GetSubKind())
}
var s3Location string
if req.AWSOIDC.IssuerS3Bucket != "" {
issuerS3URI := url.URL{
Scheme: "s3",
Host: req.AWSOIDC.IssuerS3Bucket,
Path: req.AWSOIDC.IssuerS3Prefix,
}
s3Location = issuerS3URI.String()
}
integration.SetAWSOIDCIssuerS3URI(s3Location)
integration.SetAWSOIDCRoleARN(req.AWSOIDC.RoleARN)
}
if req.OAuth != nil {
if integration.GetSubKind() != types.IntegrationSubKindGitHub {
return nil, trace.BadParameter("cannot update %q fields for a %q integration", types.IntegrationSubKindGitHub, integration.GetSubKind())
}
cred := types.PluginCredentialsV1{
Credentials: &types.PluginCredentialsV1_IdSecret{
IdSecret: &types.PluginIdSecretCredential{
Id: req.OAuth.ID,
Secret: req.OAuth.Secret,
},
},
}
if err := integration.SetCredentials(&cred); err != nil {
return nil, trace.Wrap(err)
}
}
if req.AWSRA != nil {
if integration.GetSubKind() != types.IntegrationSubKindAWSRolesAnywhere {
return nil, trace.BadParameter("cannot update %q fields for a %q integration", types.IntegrationSubKindAWSRolesAnywhere, integration.GetSubKind())
}
spec := integration.GetAWSRolesAnywhereIntegrationSpec()
spec.TrustAnchorARN = req.AWSRA.TrustAnchorARN
spec.ProfileSyncConfig = &types.AWSRolesAnywhereProfileSyncConfig{
Enabled: req.AWSRA.ProfileSyncConfig.Enabled,
ProfileARN: req.AWSRA.ProfileSyncConfig.ProfileARN,
RoleARN: req.AWSRA.ProfileSyncConfig.RoleARN,
ProfileNameFilters: req.AWSRA.ProfileSyncConfig.ProfileNameFilters,
ProfileAcceptsRoleSessionName: spec.ProfileSyncConfig.ProfileAcceptsRoleSessionName,
}
integration.SetAWSRolesAnywhereIntegrationSpec(spec)
}
if _, err := clt.UpdateIntegration(r.Context(), integration); err != nil {
return nil, trace.Wrap(err)
}
uiIg, err := ui.MakeIntegration(integration)
if err != nil {
return nil, err
}
return uiIg, nil
}
// integrationsDelete removes an Integration based on its name
func (h *Handler) integrationsDelete(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
integrationName := p.ByName("name_or_subkind")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
deleteAssociatedResources, _ := apiutils.ParseBool(r.URL.Query().Get("associatedresources"))
if _, err := clt.IntegrationsClient().DeleteIntegration(r.Context(), integrationv1.DeleteIntegrationRequest_builder{
Name: integrationName,
DeleteAssociatedResources: deleteAssociatedResources,
}.Build()); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// integrationsGet returns an Integration based on its name
func (h *Handler) integrationsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ig, err := clt.GetIntegration(r.Context(), integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
uiIg, err := ui.MakeIntegration(ig)
if err != nil {
return nil, err
}
return uiIg, nil
}
// integrationStats returns the integration stats.
func (h *Handler) integrationStats(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ig, err := clt.GetIntegration(r.Context(), integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
req := collectIntegrationStatsRequest{
logger: h.logger,
integration: ig,
discoveryConfigLister: clt.DiscoveryConfigClient(),
databaseGetter: clt,
awsOIDCClient: clt.IntegrationAWSOIDCClient(),
userTasksClient: clt.UserTasksServiceClient(),
}
summary, err := collectIntegrationStats(r.Context(), req)
if err != nil {
return nil, trace.Wrap(err)
}
return summary, nil
}
type userTasksLister interface {
ListUserTasks(ctx context.Context, pageSize int64, nextToken string, filters *usertasksv1.ListUserTasksFilters) ([]*usertasksv1.UserTask, string, error)
}
type collectIntegrationStatsRequest struct {
logger *slog.Logger
integration types.Integration
discoveryConfigLister discoveryConfigLister
databaseGetter databaseGetter
awsOIDCClient deployedDatabaseServiceLister
userTasksClient userTasksLister
}
func collectIntegrationStats(ctx context.Context, req collectIntegrationStatsRequest) (*ui.IntegrationWithSummary, error) {
ret := &ui.IntegrationWithSummary{}
uiIg, err := ui.MakeIntegration(req.integration)
if err != nil {
return nil, err
}
ret.Integration = uiIg
if req.integration != nil {
if val, ok := req.integration.GetLabel(types.CreatedByIaCLabel); ok && val == ui.IaCTerraformLabel {
ret.IsManagedByTerraform = true
}
}
tasks := allUserTasks(ctx, req.userTasksClient, usertasksv1.ListUserTasksFilters_builder{
Integration: req.integration.GetName(),
TaskState: usertasks.TaskStateOpen,
}.Build())
for task, err := range tasks {
if err != nil {
return nil, trace.Wrap(err)
}
ret.UserTasks = append(ret.UserTasks, ui.MakeUserTask(task))
switch task.GetSpec().GetTaskType() {
case usertasks.TaskTypeDiscoverEC2:
ret.AWSEC2.UnresolvedUserTasks++
case usertasks.TaskTypeDiscoverEKS:
ret.AWSEKS.UnresolvedUserTasks++
case usertasks.TaskTypeDiscoverRDS:
ret.AWSRDS.UnresolvedUserTasks++
case usertasks.TaskTypeDiscoverAzureVM:
ret.AzureVM.UnresolvedUserTasks++
}
}
ret.UnresolvedUserTasks = len(ret.UserTasks)
// Track whether any resource type is currently being scanned.
// If any are scanning, we set SyncEnd to nil after all iterations.
// TODO (avatus) might need to make this a bit more scalable in the
// future if we have a bunch of types but this is ok for now
var ec2Scanning, rdsScanning, eksScanning, azureVMScanning bool
for cfg, err := range allDiscoveryConfigs(ctx, req.discoveryConfigLister) {
if err != nil {
return nil, trace.Wrap(err)
}
summary := integrationSummaryForConfig(cfg, req.integration.GetName())
if summary == nil {
continue
}
ec2Matchers := rulesWithIntegration(cfg, types.AWSMatcherEC2, req.integration.GetName())
rdsMatchers := rulesWithIntegration(cfg, types.AWSMatcherRDS, req.integration.GetName())
eksMatchers := rulesWithIntegration(cfg, types.AWSMatcherEKS, req.integration.GetName())
azureVMMatchers := rulesWithIntegration(cfg, types.AzureMatcherVM, req.integration.GetName())
ret.AWSEC2.RulesCount += ec2Matchers
ret.AWSRDS.RulesCount += rdsMatchers
ret.AWSEKS.RulesCount += eksMatchers
ret.AzureVM.RulesCount += azureVMMatchers
if ec2Matchers != 0 {
ec2Scanning = mergeResourceTypeSummary(&ret.AWSEC2, summary.summary.GetAwsEc2(), summary.pollInterval) || ec2Scanning
}
if rdsMatchers != 0 {
rdsScanning = mergeResourceTypeSummary(&ret.AWSRDS, summary.summary.GetAwsRds(), summary.pollInterval) || rdsScanning
}
if eksMatchers != 0 {
eksScanning = mergeResourceTypeSummary(&ret.AWSEKS, summary.summary.GetAwsEks(), summary.pollInterval) || eksScanning
}
if azureVMMatchers != 0 {
azureVMScanning = mergeResourceTypeSummary(&ret.AzureVM, summary.summary.GetAzureVms(), summary.pollInterval) || azureVMScanning
}
}
// If any resource type is currently scanning, set SyncEnd to nil.
if ec2Scanning {
ret.AWSEC2.SyncEnd = nil
}
if rdsScanning {
ret.AWSRDS.SyncEnd = nil
}
if eksScanning {
ret.AWSEKS.SyncEnd = nil
}
if azureVMScanning {
ret.AzureVM.SyncEnd = nil
}
switch req.integration.GetSubKind() {
case types.IntegrationSubKindAWSRolesAnywhere:
ret.RolesAnywhereProfileSync = &ui.RolesAnywhereProfileSync{}
awsRolesAnywhereSpec := req.integration.GetAWSRolesAnywhereIntegrationSpec()
if awsRolesAnywhereSpec != nil {
ret.RolesAnywhereProfileSync.Enabled = awsRolesAnywhereSpec.ProfileSyncConfig.Enabled
}
integrationStatus := req.integration.GetStatus()
if integrationStatus.AWSRolesAnywhere != nil {
ret.RolesAnywhereProfileSync.Status = integrationStatus.AWSRolesAnywhere.LastProfileSync.Status
ret.RolesAnywhereProfileSync.ErrorMessage = integrationStatus.AWSRolesAnywhere.LastProfileSync.ErrorMessage
ret.RolesAnywhereProfileSync.SyncedProfiles = int(integrationStatus.AWSRolesAnywhere.LastProfileSync.SyncedProfiles)
ret.RolesAnywhereProfileSync.SyncStartTime = integrationStatus.AWSRolesAnywhere.LastProfileSync.StartTime
ret.RolesAnywhereProfileSync.SyncEndTime = integrationStatus.AWSRolesAnywhere.LastProfileSync.EndTime
}
case types.IntegrationSubKindAWSOIDC:
// For now, Deploying Database Services is only possible using the AWS OIDC Integration.
// When/if more integrations (eg, AWS IAM Roles Anywhere) support it, this must be updated.
ecsCount, err := countAWSOIDCDeployedDatabaseServices(ctx, req)
if err != nil {
return nil, trace.Wrap(err)
}
ret.AWSRDS.ECSDatabaseServiceCount = ecsCount
}
return ret, nil
}
func buildBriefSummaries(ctx context.Context, igs []types.Integration, uclt userTasksLister, dclt discoveryConfigLister) (map[string]*ui.BriefSummary, error) {
summaries := make(map[string]*ui.BriefSummary, len(igs))
for _, ig := range igs {
if !ig.SupportsDiscoveryResources() {
continue
}
summaries[ig.GetName()] = &ui.BriefSummary{
UnresolvedUserTasks: []ui.UserTask{},
}
}
if len(summaries) == 0 {
return summaries, nil
}
for name := range summaries {
tasks := allUserTasks(ctx, uclt, usertasksv1.ListUserTasksFilters_builder{
Integration: name,
TaskState: usertasks.TaskStateOpen,
}.Build())
for task, err := range tasks {
if err != nil {
return nil, trace.Wrap(err)
}
summaries[name].UnresolvedUserTasks = append(summaries[name].UnresolvedUserTasks, ui.MakeUserTask(task))
}
}
for cfg, err := range allDiscoveryConfigs(ctx, dclt) {
if err != nil {
return nil, trace.Wrap(err)
}
for name, rscs := range cfg.Status.IntegrationDiscoveredResources {
if _, ok := summaries[name]; !ok {
continue
}
if summaries[name].ResourcesCount == nil {
summaries[name].ResourcesCount = &ui.ResourcesCount{}
}
addResourceCounts(summaries[name].ResourcesCount, rscs.AwsEc2)
addResourceCounts(summaries[name].ResourcesCount, rscs.AwsEks)
addResourceCounts(summaries[name].ResourcesCount, rscs.AwsRds)
addResourceCounts(summaries[name].ResourcesCount, rscs.AzureVms)
}
}
return summaries, nil
}
func addResourceCounts(rc *ui.ResourcesCount, dr *discoveryconfigv1.ResourcesDiscoveredSummary) {
if rc == nil || dr == nil {
return
}
rc.Found += int(dr.GetFound())
rc.Enrolled += int(dr.GetEnrolled())
rc.Failed += int(dr.GetFailed())
}
func allUserTasks(
ctx context.Context,
lister userTasksLister,
filters *usertasksv1.ListUserTasksFilters,
) iter.Seq2[*usertasksv1.UserTask, error] {
return clientutils.Resources(ctx,
func(ctx context.Context, pageSize int, nextToken string) ([]*usertasksv1.UserTask, string, error) {
return lister.ListUserTasks(ctx, int64(pageSize), nextToken, filters)
})
}
func allDiscoveryConfigs(
ctx context.Context,
lister discoveryConfigLister,
) iter.Seq2[*discoveryconfig.DiscoveryConfig, error] {
return clientutils.Resources(ctx,
func(ctx context.Context, pageSize int, nextToken string) ([]*discoveryconfig.DiscoveryConfig, string, error) {
return lister.ListDiscoveryConfigs(ctx, pageSize, nextToken)
})
}
func countAWSOIDCDeployedDatabaseServices(ctx context.Context, req collectIntegrationStatsRequest) (int, error) {
regions, err := fetchRelevantAWSRegions(ctx, req.databaseGetter, req.discoveryConfigLister)
if err != nil {
return 0, trace.Wrap(err)
}
services, err := listDeployedDatabaseServices(ctx, req.logger, req.integration.GetName(), regions, req.awsOIDCClient)
if err != nil {
// The number of ECS Database Services is shown when listing the integration status.
// However, listing ECS Services is only possible after the user goes through the RDS enrollment flows, which adds the required policy to the IAM Role.
// If this calls returns an access denied, we assume the user doesn't have the required IAM Policies in their IAM Role and show 0 instead.
if trace.IsAccessDenied(err) {
return 0, nil
}
return 0, trace.Wrap(err)
}
return len(services), nil
}
type integrationSummaryWithPollInterval struct {
summary *discoveryconfigv1.DiscoverSummary
pollInterval time.Duration
}
// integrationSummaryForConfig returns the most recently updated server's summary
// for the given integration. In HA deployments, multiple Discovery Services may
// report summaries for the same resources, so we use only the most recent one
// to avoid double-counting.
func integrationSummaryForConfig(cfg *discoveryconfig.DiscoveryConfig, integrationName string) *integrationSummaryWithPollInterval {
var mostRecent *integrationSummaryWithPollInterval
var mostRecentTime time.Time
for _, serverStatus := range cfg.Status.ServerStatus {
if serverStatus == nil || serverStatus.DiscoveryStatusServer == nil {
continue
}
discoverSummary, ok := serverStatus.GetIntegrationSummaries()[integrationName]
if !ok {
continue
}
lastUpdate := serverStatus.GetLastUpdate().AsTime()
if mostRecent == nil || lastUpdate.After(mostRecentTime) {
mostRecent = &integrationSummaryWithPollInterval{
summary: discoverSummary,
pollInterval: serverStatus.GetPollInterval().AsDuration(),
}
mostRecentTime = lastUpdate
}
}
return mostRecent
}
// mergeResourceTypeSummary merges resource summary data into the aggregated summary.
// It returns true if the resource is currently being scanned (has no SyncEnd time yet).
func mergeResourceTypeSummary(in *ui.ResourceTypeSummary, resourceSummary *discoveryconfigv1.ResourceSummary, pollInterval time.Duration) bool {
if resourceSummary == nil {
return false
}
previous := resourceSummary.GetPrevious()
if previous != nil {
in.ResourcesFound += int(previous.GetFound())
in.ResourcesEnrollmentSuccess += int(previous.GetEnrolled())
in.ResourcesEnrollmentFailed += int(previous.GetFailed())
in.DiscoverLastSync = latestTime(in.DiscoverLastSync, previous.GetSyncEnd().AsTime())
}
isScanning := false
syncEndUpdated := false
current := resourceSummary.GetCurrent()
if current != nil {
in.SyncStart = latestTime(in.SyncStart, current.GetSyncStart().AsTime())
if current.GetSyncEnd().AsTime().IsZero() {
isScanning = true
} else {
prevSyncEnd := in.SyncEnd
in.SyncEnd = latestTime(in.SyncEnd, current.GetSyncEnd().AsTime())
syncEndUpdated = in.SyncEnd != prevSyncEnd
}
} else if previous != nil {
in.SyncStart = latestTime(in.SyncStart, previous.GetSyncStart().AsTime())
prevSyncEnd := in.SyncEnd
in.SyncEnd = latestTime(in.SyncEnd, previous.GetSyncEnd().AsTime())
syncEndUpdated = in.SyncEnd != prevSyncEnd
}
if pollInterval > 0 && (in.PollIntervalSeconds == 0 || syncEndUpdated) {
in.PollIntervalSeconds = int(pollInterval.Seconds())
}
return isScanning
}
func latestTime(current *time.Time, new time.Time) *time.Time {
if new.IsZero() {
return current
}
if current == nil {
return &new
}
if current.Before(new) {
return &new
}
return current
}
// rulesWithIntegration returns the number of Rules for a given integration and matcher type in the DiscoveryConfig.
// A Rule is similar to a DiscoveryConfig's Matcher, eg DiscoveryConfig.Spec.AWS.[<Matcher>], however, a Rule has a single region.
// This means that the number of Rules for a given Matcher is equal to the number of regions on that Matcher.
func rulesWithIntegration(dc *discoveryconfig.DiscoveryConfig, matcherType string, integration string) int {
ret := 0
for _, matcher := range dc.Spec.AWS {
if matcher.Integration != integration {
continue
}
if !slices.Contains(matcher.Types, matcherType) {
continue
}
ret += len(matcher.Regions)
}
for _, matcher := range dc.Spec.Azure {
if matcher.Integration != integration {
continue
}
if !slices.Contains(matcher.Types, matcherType) {
continue
}
ret += len(matcher.Regions)
}
return ret
}
// integrationDiscoveryRules returns the Discovery Rules that are using a given integration.
// A Discovery Rule is just like a DiscoveryConfig Matcher, except that it breaks down by region.
// So, if a Matcher exists for two regions, that will be represented as two Rules.
// Accepts the following query params:
// startKey: indicator for pagination, should be the value of the last reponse's `nextItem`, or absent for a the starting page
// resourceType: which resource type to return, one of ec2, eks, rds
// regions: only rules for regions listed are returned (omit query to include all regions)
func (h *Handler) integrationDiscoveryRules(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
values := r.URL.Query()
startKey := values.Get("startKey")
resourceType := values.Get("resourceType")
regionsFilter := values["regions"]
// the regions key is always sent as a query param but is not always populated (®ions=)
// this results in a slice containing a single empty string
if len(regionsFilter) == 1 && regionsFilter[0] == "" {
regionsFilter = nil
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ig, err := clt.GetIntegration(r.Context(), integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
rules, err := collectAutoDiscoveryRules(r.Context(), ig.GetName(), startKey, resourceType, regionsFilter, clt.DiscoveryConfigClient())
if err != nil {
return nil, trace.Wrap(err)
}
return rules, nil
}
// collectAutoDiscoveryRules will iterate over all DiscoveryConfigs's Matchers and collect the Discovery Rules that exist in them for the given integration.
// It can also be filtered by Matcher Type (eg ec2, rds, eks) and a regionsFilter list (eg, us-east-1, us-east-2)
// A Discovery Rule is a close match to a DiscoveryConfig's Matcher, except that it will count as many rules as regions exist.
// Eg if a DiscoveryConfig's Matcher has two regions, then it will output two (almost equal) Rules, one for each Region.
func collectAutoDiscoveryRules(
ctx context.Context,
integrationName string,
nextPage string,
resourceTypeFilter string,
regionsFilter []string,
clt interface {
ListDiscoveryConfigs(ctx context.Context, pageSize int, nextToken string) ([]*discoveryconfig.DiscoveryConfig, string, error)
},
) (ui.IntegrationDiscoveryRules, error) {
const (
maxPerPage = 100
)
var ret ui.IntegrationDiscoveryRules
for {
discoveryConfigs, nextToken, err := clt.ListDiscoveryConfigs(ctx, 0, nextPage)
if err != nil {
return ret, trace.Wrap(err)
}
for _, dc := range discoveryConfigs {
ret.Rules = append(ret.Rules,
collectAutoDiscoveryRulesFromDiscoveryConfig(dc, integrationName, resourceTypeFilter, regionsFilter)...,
)
}
ret.NextKey = nextToken
if nextToken == "" || len(ret.Rules) > maxPerPage {
break
}
nextPage = nextToken
}
return ret, nil
}
func collectAutoDiscoveryRulesFromDiscoveryConfig(dc *discoveryconfig.DiscoveryConfig, integrationName, resourceTypeFilter string, regionsFilter []string) []ui.IntegrationDiscoveryRule {
lastSync := &dc.Status.LastSyncTime
if lastSync.IsZero() {
lastSync = nil
}
awsRules := collectAWSAutoDiscoveryRulesFromDiscoveryConfig(dc, integrationName, resourceTypeFilter, regionsFilter, lastSync)
azureRules := collectAzureAutoDiscoveryRulesFromDiscoveryConfig(dc, integrationName, resourceTypeFilter, regionsFilter, lastSync)
return append(awsRules, azureRules...)
}
func collectAWSAutoDiscoveryRulesFromDiscoveryConfig(dc *discoveryconfig.DiscoveryConfig, integrationName, resourceTypeFilter string, regionsFilter []string, lastSync *time.Time) (ret []ui.IntegrationDiscoveryRule) {
for _, matcher := range dc.Spec.AWS {
if matcher.Integration != integrationName {
continue
}
for _, resourceType := range matcher.Types {
if resourceTypeFilter != "" && resourceType != resourceTypeFilter {
continue
}
for _, region := range matcher.Regions {
if len(regionsFilter) > 0 && !slices.Contains(regionsFilter, region) {
continue
}
uiLabels := make([]libui.Label, 0, len(matcher.Tags))
for labelKey, labelValues := range matcher.Tags {
for _, labelValue := range labelValues {
uiLabels = append(uiLabels, libui.Label{
Name: labelKey,
Value: labelValue,
})
}
}
rule := ui.IntegrationDiscoveryRule{
ResourceType: resourceType,
Region: region,
LabelMatcher: uiLabels,
DiscoveryConfig: dc.GetName(),
LastSync: lastSync,
}
if resourceType == "eks" {
kubeAppDiscovery := matcher.KubeAppDiscovery
rule.KubeAppDiscovery = &kubeAppDiscovery
}
ret = append(ret, rule)
}
}
}
return
}
func collectAzureAutoDiscoveryRulesFromDiscoveryConfig(dc *discoveryconfig.DiscoveryConfig, integrationName, resourceTypeFilter string, regionsFilter []string, lastSync *time.Time) (ret []ui.IntegrationDiscoveryRule) {
for _, matcher := range dc.Spec.Azure {
if matcher.Integration != integrationName {
continue
}
for _, resourceType := range matcher.Types {
if resourceTypeFilter != "" && resourceType != resourceTypeFilter {
continue
}
for _, region := range matcher.Regions {
if len(regionsFilter) > 0 && !slices.Contains(regionsFilter, region) {
continue
}
uiLabels := make([]libui.Label, 0, len(matcher.ResourceTags))
for labelKey, labelValues := range matcher.ResourceTags {
for _, labelValue := range labelValues {
uiLabels = append(uiLabels, libui.Label{
Name: labelKey,
Value: labelValue,
})
}
}
ret = append(ret, ui.IntegrationDiscoveryRule{
ResourceType: resourceType,
Region: region,
LabelMatcher: uiLabels,
Subscriptions: matcher.Subscriptions,
ResourceGroups: matcher.ResourceGroups,
DiscoveryConfig: dc.GetName(),
LastSync: lastSync,
})
}
}
}
return
}
// integrationsList returns a page of Integrations
func (h *Handler) integrationsList(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
startKey := values.Get("startKey")
igs, nextKey, err := clt.ListIntegrations(r.Context(), int(limit), startKey)
if err != nil {
return nil, trace.Wrap(err)
}
var summaries map[string]*ui.BriefSummary
withSummaries, err := parseBoolWithDefault(values.Get("withSummaries"), false)
if err != nil {
return nil, trace.Wrap(err)
}
if withSummaries {
summaries, err = buildBriefSummaries(r.Context(), igs, clt.UserTasksServiceClient(), clt.DiscoveryConfigClient())
if err != nil {
return nil, trace.Wrap(err)
}
}
items, err := ui.MakeIntegrations(igs)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.IntegrationsListResponse{
Items: items,
NextKey: nextKey,
Summaries: summaries,
}, nil
}
// integrationsMsTeamsAppZipGet generates and returns the app.zip required for the MsTeams plugin with the given name.
func (h *Handler) integrationsMsTeamsAppZipGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
plugin, err := clt.PluginsClient().GetPlugin(r.Context(), pluginspb.GetPluginRequest_builder{
Name: p.ByName("plugin"),
WithSecrets: false,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
spec, ok := plugin.Spec.Settings.(*types.PluginSpecV1_Msteams)
if !ok {
return nil, trace.BadParameter("plugin specified was not of type MsTeams")
}
w.Header().Add("Content-Type", "application/zip")
w.Header().Add("Content-Disposition", "attachment; filename=app.zip")
err = msteams.WriteAppZipTo(w, msteams.ConfigTemplatePayload{
AppID: spec.Msteams.AppId,
TenantID: spec.Msteams.TenantId,
TeamsAppID: spec.Msteams.TeamsAppId,
})
if err != nil {
return nil, trace.Wrap(err)
}
return nil, nil
}
func (h *Handler) integrationsExportCA(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := clt.IntegrationsClient().ExportIntegrationCertAuthorities(r.Context(), integrationv1.ExportIntegrationCertAuthoritiesRequest_builder{
Integration: integrationName,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
uiCAKeySet, err := ui.MakeCAKeySet(resp.GetCertAuthorities())
return uiCAKeySet, trace.Wrap(err)
}
/*
* 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"
"fmt"
"log/slog"
"maps"
"net/http"
"slices"
"strconv"
"strings"
"github.com/aws/aws-sdk-go-v2/aws/arn"
"github.com/coreos/go-semver/semver"
"github.com/google/safetext/shsprintf"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"google.golang.org/grpc"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
apidefaults "github.com/gravitational/teleport/api/defaults"
integrationv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/integration/v1"
presencev1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/presence/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/discoveryconfig"
"github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/lib/auth/authclient"
autoupdateversion "github.com/gravitational/teleport/lib/automaticupgrades/version"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/integrations/awsoidc"
"github.com/gravitational/teleport/lib/integrations/awsoidc/deployserviceconfig"
kubeutils "github.com/gravitational/teleport/lib/kube/utils"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
libui "github.com/gravitational/teleport/lib/ui"
libutils "github.com/gravitational/teleport/lib/utils"
awsutils "github.com/gravitational/teleport/lib/utils/aws"
"github.com/gravitational/teleport/lib/utils/oidc"
"github.com/gravitational/teleport/lib/utils/set"
"github.com/gravitational/teleport/lib/web/scripts/oneoff"
"github.com/gravitational/teleport/lib/web/ui"
)
// awsOIDCListDatabases returns a list of databases using the ListDatabases action of the AWS OIDC Integration.
func (h *Handler) awsOIDCListDatabases(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCListDatabasesRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
listDatabasesResp, err := clt.IntegrationAWSOIDCClient().ListDatabases(ctx, integrationv1.ListDatabasesRequest_builder{
Integration: integrationName,
Region: req.Region,
RdsType: req.RDSType,
Engines: req.Engines,
NextToken: req.NextToken,
VpcId: req.VPCID,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSOIDCListDatabasesResponse{
NextToken: listDatabasesResp.GetNextToken(),
Databases: ui.MakeDatabases(listDatabasesResp.GetDatabases(), accessChecker, h.cfg.DatabaseREPLRegistry),
}, nil
}
// awsOIDCDeployService deploys a Discovery Service and a Database Service in Amazon ECS.
func (h *Handler) awsOIDCDeployService(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCDeployServiceRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
databaseAgentMatcherLabels := make(types.Labels, len(req.DatabaseAgentMatcherLabels)+3)
for _, label := range req.DatabaseAgentMatcherLabels {
databaseAgentMatcherLabels[label.Name] = utils.Strings{label.Value}
}
// DELETE in 19.0: delete only the outer if block (checking labels == 0).
// The outer block is required since older UI's will not
// send these values to the backend, but instead send custom labels (the UI
// will require at least one label before proceeding).
// Newer UI's will not send any labels, but instead send the required
// fields for default labels.
if len(req.DatabaseAgentMatcherLabels) == 0 {
if req.VPCID == "" {
return nil, trace.BadParameter("vpc ID is required")
}
if req.Region == "" {
return nil, trace.BadParameter("AWS region is required")
}
if req.AccountID == "" {
return nil, trace.BadParameter("AWS account ID is required")
}
// Add default labels.
databaseAgentMatcherLabels[types.DiscoveryLabelVPCID] = []string{req.VPCID}
databaseAgentMatcherLabels[types.DiscoveryLabelRegion] = []string{req.Region}
databaseAgentMatcherLabels[types.DiscoveryLabelAccountID] = []string{req.AccountID}
}
iamTokenName := deployserviceconfig.DefaultTeleportIAMTokenName
teleportConfigString, err := deployserviceconfig.GenerateTeleportConfigString(
h.PublicProxyAddr(),
iamTokenName,
databaseAgentMatcherLabels,
)
if err != nil {
return nil, trace.Wrap(err)
}
teleportVersionTag := teleport.Version
if automaticUpgrades(h.GetClusterFeatures()) {
const group, updaterUUID = "", ""
autoUpdateVersion, err := h.autoUpdateResolver.GetVersion(r.Context(), group, updaterUUID)
if err != nil {
h.logger.WarnContext(r.Context(),
"Cannot read autoupdate target version, falling back to our own version",
"error", err,
"version", teleport.Version)
} else {
teleportVersionTag = autoUpdateVersion.String()
}
}
deployServiceResp, err := clt.IntegrationAWSOIDCClient().DeployService(ctx, integrationv1.DeployServiceRequest_builder{
DeploymentJoinTokenName: iamTokenName,
DeploymentMode: req.DeploymentMode,
TeleportConfigString: teleportConfigString,
Integration: integrationName,
Region: req.Region,
SecurityGroups: req.SecurityGroups,
SubnetIds: req.SubnetIDs,
TaskRoleArn: req.TaskRoleARN,
TeleportVersion: teleportVersionTag,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSOIDCDeployServiceResponse{
ClusterARN: deployServiceResp.GetClusterArn(),
ServiceARN: deployServiceResp.GetServiceArn(),
TaskDefinitionARN: deployServiceResp.GetTaskDefinitionArn(),
ServiceDashboardURL: deployServiceResp.GetServiceDashboardUrl(),
}, nil
}
// awsOIDCDeployDatabaseService deploys a Database Service in Amazon ECS.
func (h *Handler) awsOIDCDeployDatabaseServices(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCDeployDatabaseServiceRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
teleportVersionTag := teleport.Version
if automaticUpgrades(h.GetClusterFeatures()) {
const group, updaterUUID = "", ""
autoUpdateVersion, err := h.autoUpdateResolver.GetVersion(r.Context(), group, updaterUUID)
if err != nil {
h.logger.WarnContext(r.Context(),
"Cannot read autoupdate target version, falling back to self version.",
"error", err,
"version", teleport.Version)
} else {
teleportVersionTag = autoUpdateVersion.String()
}
}
iamTokenName := deployserviceconfig.DefaultTeleportIAMTokenName
deployments := make([]*integrationv1.DeployDatabaseServiceDeployment, 0, len(req.Deployments))
for _, d := range req.Deployments {
teleportConfigString, err := deployserviceconfig.GenerateTeleportConfigString(
h.PublicProxyAddr(),
iamTokenName,
types.Labels{
types.DiscoveryLabelVPCID: []string{d.VPCID},
types.DiscoveryLabelRegion: []string{req.Region},
types.DiscoveryLabelAccountID: []string{req.AccountID},
},
)
if err != nil {
return nil, trace.Wrap(err)
}
deployments = append(deployments, integrationv1.DeployDatabaseServiceDeployment_builder{
VpcId: d.VPCID,
SubnetIds: d.SubnetIDs,
SecurityGroups: d.SecurityGroups,
TeleportConfigString: teleportConfigString,
}.Build())
}
deployServiceResp, err := clt.IntegrationAWSOIDCClient().DeployDatabaseService(ctx, integrationv1.DeployDatabaseServiceRequest_builder{
Integration: integrationName,
Region: req.Region,
TaskRoleArn: req.TaskRoleARN,
Deployments: deployments,
TeleportVersion: teleportVersionTag,
DeploymentJoinTokenName: deployserviceconfig.DefaultTeleportIAMTokenName,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSOIDCDeployDatabaseServiceResponse{
ClusterARN: deployServiceResp.GetClusterArn(),
ClusterDashboardURL: deployServiceResp.GetClusterDashboardUrl(),
}, nil
}
// awsOIDCListDeployedDatabaseService lists the deployed Database Services in Amazon ECS.
func (h *Handler) awsOIDCListDeployedDatabaseService(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
regions, err := regionsForListingDeployedDatabaseService(ctx, r, clt, clt.DiscoveryConfigClient())
if err != nil {
return nil, trace.Wrap(err)
}
if len(regions) == 0 {
// return an empty list if there are no relevant regions in which to fetch database services
return ui.AWSOIDCListDeployedDatabaseServiceResponse{}, nil
}
s, err := listDeployedDatabaseServices(ctx, h.logger, integrationName, regions, clt.IntegrationAWSOIDCClient())
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSOIDCListDeployedDatabaseServiceResponse{
Services: s,
}, nil
}
func extractAWSRegionsFromQuery(r *http.Request) ([]string, error) {
var ret []string
for _, region := range r.URL.Query()["regions"] {
if region == "" {
// no regions passed in params, empty key
return ret, nil
}
if err := aws.IsValidRegion(region); err != nil {
return nil, trace.BadParameter("invalid region %s", region)
}
ret = append(ret, region)
}
return ret, nil
}
// regionsForListingDeployedDatabaseService fetches relevant AWS regions and parses the regions query param.
// If no query params are present, relevant regions are returned.
// If query params are present, we take the intersection of relevant regions and filter regions to avoid requesting
// services which have not been set up which would result in an error.
// ex: relevant = ["us-west-1"]; params = []; returns ["us-west-1"]
// ex: relevant = []; params = ["us-west-1"]; returns []
// ex: relevant = ["us-west-1"]; params = ["us-west-1"]; returns ["us-west-1"]
// ex: relevant = ["us-west-1"]; params = ["us-west-2"]; returns []
func regionsForListingDeployedDatabaseService(ctx context.Context, r *http.Request, authClient databaseGetter, discoveryConfigsClient discoveryConfigLister) ([]string, error) {
// use the auth client & discover configs to collect a list of relevant AWS regions
relevant, err := fetchRelevantAWSRegions(ctx, authClient, discoveryConfigsClient)
if err != nil {
return nil, trace.Wrap(err)
}
if r.URL.Query().Has("regions") {
params, err := extractAWSRegionsFromQuery(r)
if err != nil {
return nil, trace.Wrap(err)
}
if len(params) > 0 {
a := set.New(relevant...)
b := set.New(params...)
a.Intersection(b)
return a.Elements(), nil
}
}
return relevant, nil
}
type databaseGetter interface {
GetResources(ctx context.Context, req *proto.ListResourcesRequest) (*proto.ListResourcesResponse, error)
GetDatabases(context.Context) ([]types.Database, error)
}
type discoveryConfigLister interface {
ListDiscoveryConfigs(ctx context.Context, pageSize int, nextToken string) ([]*discoveryconfig.DiscoveryConfig, string, error)
}
func fetchRelevantAWSRegions(ctx context.Context, authClient databaseGetter, discoveryConfigsClient discoveryConfigLister) ([]string, error) {
regionsSet := make(map[string]struct{})
// Collect Regions from Database resources.
databases, err := authClient.GetDatabases(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
for _, resource := range databases {
regionsSet[resource.GetAWS().Region] = struct{}{}
regionsSet[resource.GetAllLabels()[types.DiscoveryLabelRegion]] = struct{}{}
}
// Iterate over all DatabaseServices and fetch their AWS Region in the matchers.
var nextPageKey string
for {
req := &proto.ListResourcesRequest{
ResourceType: types.KindDatabaseService,
Limit: defaults.MaxIterationLimit,
StartKey: nextPageKey,
Labels: map[string]string{types.AWSOIDCAgentLabel: types.True},
}
page, err := client.GetResourcePage[types.DatabaseService](ctx, authClient, req)
if err != nil {
return nil, trace.Wrap(err)
}
maps.Copy(regionsSet, extractRegionsFromDatabaseServicesPage(page.Resources))
if page.NextKey == "" {
break
}
nextPageKey = page.NextKey
}
// Iterate over all DiscoveryConfigs and fetch their AWS Region in AWS Matchers.
nextPageKey = ""
for {
resp, respNextPageKey, err := discoveryConfigsClient.ListDiscoveryConfigs(ctx, defaults.MaxIterationLimit, nextPageKey)
if err != nil {
return nil, trace.Wrap(err)
}
maps.Copy(regionsSet, extractRegionsFromDiscoveryConfigPage(resp))
if respNextPageKey == "" {
break
}
nextPageKey = respNextPageKey
}
// Drop any invalid region.
ret := make([]string, 0, len(regionsSet))
for region := range regionsSet {
if aws.IsValidRegion(region) == nil {
ret = append(ret, region)
}
}
return ret, nil
}
func extractRegionsFromDatabaseServicesPage(dbServices []types.DatabaseService) map[string]struct{} {
regionsSet := make(map[string]struct{})
for _, resource := range dbServices {
for _, matcher := range resource.GetResourceMatchers() {
if matcher.Labels == nil {
continue
}
for labelKey, labelValues := range *matcher.Labels {
if labelKey != types.DiscoveryLabelRegion {
continue
}
for _, labelValue := range labelValues {
regionsSet[labelValue] = struct{}{}
}
}
}
}
return regionsSet
}
func extractRegionsFromDiscoveryConfigPage(discoveryConfigs []*discoveryconfig.DiscoveryConfig) map[string]struct{} {
regionsSet := make(map[string]struct{})
for _, dc := range discoveryConfigs {
for _, awsMatcher := range dc.Spec.AWS {
for _, region := range awsMatcher.Regions {
regionsSet[region] = struct{}{}
}
}
}
return regionsSet
}
type deployedDatabaseServiceLister interface {
ListDeployedDatabaseServices(ctx context.Context, in *integrationv1.ListDeployedDatabaseServicesRequest, opts ...grpc.CallOption) (*integrationv1.ListDeployedDatabaseServicesResponse, error)
}
func listDeployedDatabaseServices(ctx context.Context,
logger *slog.Logger,
integrationName string,
regions []string,
awsOIDCClient deployedDatabaseServiceLister,
) ([]ui.AWSOIDCDeployedDatabaseService, error) {
var services []ui.AWSOIDCDeployedDatabaseService
for _, region := range regions {
var nextToken string
for {
resp, err := awsOIDCClient.ListDeployedDatabaseServices(ctx, integrationv1.ListDeployedDatabaseServicesRequest_builder{
Integration: integrationName,
Region: region,
NextToken: nextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
for _, deployedDatabaseService := range resp.GetDeployedDatabaseServices() {
matchingLabels, err := matchingLabelsFromDeployedService(deployedDatabaseService)
if err != nil {
logger.WarnContext(ctx, "Failed to obtain teleport config string from ECS Service",
"ecs_service", deployedDatabaseService.GetServiceDashboardUrl(),
"error", err,
)
}
validTeleportConfigFound := err == nil
services = append(services, ui.AWSOIDCDeployedDatabaseService{
Name: deployedDatabaseService.GetName(),
DashboardURL: deployedDatabaseService.GetServiceDashboardUrl(),
MatchingLabels: matchingLabels,
ValidTeleportConfig: validTeleportConfigFound,
})
}
if resp.GetNextToken() == "" {
break
}
nextToken = resp.GetNextToken()
}
}
return services, nil
}
func matchingLabelsFromDeployedService(deployedDatabaseService *integrationv1.DeployedDatabaseService) ([]libui.Label, error) {
commandArgs := deployedDatabaseService.GetContainerCommand()
// This command is what starts the teleport agent in the ECS Service Fargate container.
// See deployservice.go/upsertTask for details.
// It is expected to have at least 3 values, even if dumb-init is removed in the future.
if len(commandArgs) < 3 {
return nil, trace.BadParameter("unexpected command size, expected at least 3 args, got %d", len(commandArgs))
}
// The command should have a --config-string flag and then the teleport's base64 encoded configuration as argument
teleportConfigStringFlagIdx := slices.Index(commandArgs, "--config-string")
if teleportConfigStringFlagIdx == -1 {
return nil, trace.BadParameter("missing --config-string flag in container command")
}
if len(commandArgs) < teleportConfigStringFlagIdx+1 {
return nil, trace.BadParameter("missing --config-string argument in container command")
}
teleportConfigString := commandArgs[teleportConfigStringFlagIdx+1]
labelMatchers, err := deployserviceconfig.ParseResourceLabelMatchers(teleportConfigString)
if err != nil {
return nil, trace.Wrap(err)
}
var matchingLabels []libui.Label
for labelKey, labelValues := range labelMatchers {
for _, labelValue := range labelValues {
matchingLabels = append(matchingLabels, libui.Label{
Name: labelKey,
Value: labelValue,
})
}
}
return matchingLabels, nil
}
// awsOIDCConfigureDeployServiceIAM returns a script that configures the required IAM permissions to enable the usage of DeployService action.
func (h *Handler) awsOIDCConfigureDeployServiceIAM(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
ctx := r.Context()
queryParams := r.URL.Query()
clusterName, err := h.GetProxyClient().GetDomainName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
integrationName := queryParams.Get("integrationName")
if len(integrationName) == 0 {
return nil, trace.BadParameter("missing integrationName param")
}
// Ensure the IntegrationName is valid.
_, err = h.GetProxyClient().GetIntegration(ctx, integrationName)
// NotFound error is ignored to prevent disclosure of whether the integration exists in a public/no-auth endpoint.
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
awsRegion := queryParams.Get("awsRegion")
if err := aws.IsValidRegion(awsRegion); err != nil {
return nil, trace.BadParameter("invalid awsRegion")
}
awsAccountID := queryParams.Get("awsAccountID")
if err := aws.IsValidAccountID(awsAccountID); err != nil {
return nil, trace.Wrap(err, "invalid awsAccountID")
}
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
taskRole := queryParams.Get("taskRole")
if err := aws.IsValidIAMRoleName(taskRole); err != nil {
return nil, trace.BadParameter("invalid taskRole")
}
// The script must execute the following command:
// teleport integration configure deployservice-iam
argsList := []string{
"integration", "configure", "deployservice-iam",
fmt.Sprintf("--cluster=%s", shsprintf.EscapeDefaultContext(clusterName)),
fmt.Sprintf("--name=%s", shsprintf.EscapeDefaultContext(integrationName)),
fmt.Sprintf("--aws-region=%s", shsprintf.EscapeDefaultContext(awsRegion)),
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--task-role=%s", shsprintf.EscapeDefaultContext(taskRole)),
fmt.Sprintf("--aws-account-id=%s", shsprintf.EscapeDefaultContext(awsAccountID)),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to complete the database enrollment.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// awsOIDCConfigureAppAccessIAM returns a script that configures the required IAM permissions to enable App Access
// using the AWS OIDC Credentials.
// Only IAM Roles with `teleport.dev/integration: Allowed` Tag can be used.
// It receives the IAM Role from a query param "role".
// The script is returned using the Content-Type "text/x-shellscript". No Content-Disposition header is set.
func (h *Handler) awsOIDCConfigureAWSAppAccessIAM(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
queryParams := r.URL.Query()
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
// The script must execute the following command:
// teleport integration configure aws-app-access
argsList := []string{
"integration", "configure", "aws-app-access-iam",
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to use AWS App Access.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// awsOIDCConfigureEC2SSMIAM returns a script that configures AWS IAM Policies and creates an SSM Document
// to enable EC2 Auto Discover Script mode, using the AWS OIDC Credentials.
// It receives the IAM Role, AWS Region and SSM Document Name from query params ("role", "awsRegion" and "ssmDocument").
//
// The script is returned using the Content-Type "text/x-shellscript".
// No Content-Disposition header is set.
func (h *Handler) awsOIDCConfigureEC2SSMIAM(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
queryParams := r.URL.Query()
integrationName := queryParams.Get("integrationName")
if len(integrationName) == 0 {
return nil, trace.BadParameter("missing integrationName param")
}
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
region := queryParams.Get("awsRegion")
if err := aws.IsValidRegion(region); err != nil {
return nil, trace.BadParameter("invalid region %q", region)
}
awsAccountID := queryParams.Get("awsAccountID")
if err := aws.IsValidAccountID(awsAccountID); err != nil {
return nil, trace.Wrap(err, "invalid awsAccountID")
}
// SSM Document is not required, since the user might want to use a pre-defined one. Eg, AWS-RunShellScript.
ssmDocumentName := queryParams.Get("ssmDocument")
// PublicProxyAddr() might return tenant.teleport.sh
// However, the expected format for --proxy-public-url includes the protocol `https://`
proxyPublicURL := h.PublicProxyAddr()
if !strings.HasPrefix(proxyPublicURL, "https://") {
proxyPublicURL = "https://" + proxyPublicURL
}
clusterName, err := h.GetProxyClient().GetDomainName(r.Context())
if err != nil {
return nil, trace.Wrap(err)
}
// The script must execute the following command:
// teleport integration configure ec2-ssm-iam
argsList := []string{
"integration", "configure", "ec2-ssm-iam",
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--aws-region=%s", shsprintf.EscapeDefaultContext(region)),
fmt.Sprintf("--ssm-document-name=%s", shsprintf.EscapeDefaultContext(ssmDocumentName)),
fmt.Sprintf("--proxy-public-url=%s", shsprintf.EscapeDefaultContext(proxyPublicURL)),
fmt.Sprintf("--cluster=%s", shsprintf.EscapeDefaultContext(clusterName)),
fmt.Sprintf("--name=%s", shsprintf.EscapeDefaultContext(integrationName)),
fmt.Sprintf("--aws-account-id=%s", shsprintf.EscapeDefaultContext(awsAccountID)),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to finish the EC2 auto discover set up.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// awsOIDCConfigureEKSIAM returns a script that configures the required IAM permissions to enroll EKS clusters into Teleport.
func (h *Handler) awsOIDCConfigureEKSIAM(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
queryParams := r.URL.Query()
awsRegion := queryParams.Get("awsRegion")
if err := aws.IsValidRegion(awsRegion); err != nil {
return nil, trace.BadParameter("invalid aws region")
}
awsAccountID := queryParams.Get("awsAccountID")
if err := aws.IsValidAccountID(awsAccountID); err != nil {
return nil, trace.Wrap(err, "invalid awsAccountID")
}
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
// The script must execute the following command:
// "teleport integration configure eks-iam"
argsList := []string{
"integration", "configure", "eks-iam",
fmt.Sprintf("--aws-region=%s", shsprintf.EscapeDefaultContext(awsRegion)),
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--aws-account-id=%s", shsprintf.EscapeDefaultContext(awsAccountID)),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to complete the EKS enrollment.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// handlerVersionGetter implements version.Getter by wrapping the Handler's autoupdate resolver
// and falling back to teleport.Version (the proxy's own version) when no autoupdate target is configured.
type handlerVersionGetter struct {
*Handler
}
// GetVersion implements version.Getter.
func (h *handlerVersionGetter) GetVersion(ctx context.Context) (*semver.Version, error) {
const group, updaterUUID = "", ""
v, err := h.autoUpdateResolver.GetVersion(ctx, group, updaterUUID)
if err == nil {
return v, nil
}
var noNewVersionErr *autoupdateversion.NoNewVersionError
if !errors.As(trace.Unwrap(err), &noNewVersionErr) {
return nil, trace.Wrap(err)
}
return autoupdateversion.EnsureSemver(teleport.Version)
}
// awsOIDCEnrollEKSClusters enroll EKS clusters by installing teleport-kube-agent Helm chart on them.
// v2 endpoint introduces "extraLabels" field.
func (h *Handler) awsOIDCEnrollEKSClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCEnrollEKSClustersRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
versionGetter := &handlerVersionGetter{h}
agentVersion, err := kubeutils.GetKubeAgentVersion(ctx, versionGetter)
if err != nil {
return nil, trace.Wrap(err)
}
extraLabels := make(map[string]string, len(req.ExtraLabels))
for _, label := range req.ExtraLabels {
extraLabels[label.Name] = label.Value
}
response, err := clt.IntegrationAWSOIDCClient().EnrollEKSClusters(ctx, integrationv1.EnrollEKSClustersRequest_builder{
Integration: integrationName,
Region: req.Region,
EksClusterNames: req.ClusterNames,
EnableAppDiscovery: req.EnableAppDiscovery,
AgentVersion: agentVersion.String(),
ExtraLabels: extraLabels,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
var data []ui.EKSClusterEnrollmentResult
for _, result := range response.GetResults() {
data = append(data, ui.EKSClusterEnrollmentResult{
ClusterName: result.GetEksClusterName(),
Error: result.GetError(),
ResourceId: result.GetResourceId(),
},
)
}
return ui.AWSOIDCEnrollEKSClustersResponse{
Results: data,
}, nil
}
// awsOIDCListEKSClusters returns a list of EKS clusters using the ListEKSClusters action of the AWS OIDC integration.
func (h *Handler) awsOIDCListEKSClusters(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCListEKSClustersRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
listResp, err := clt.IntegrationAWSOIDCClient().ListEKSClusters(ctx, integrationv1.ListEKSClustersRequest_builder{
Integration: integrationName,
Region: req.Region,
NextToken: req.NextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSOIDCListEKSClustersResponse{
NextToken: listResp.GetNextToken(),
Clusters: ui.MakeEKSClusters(listResp.GetClusters()),
}, nil
}
// awsOIDCListSecurityGroups returns a list of VPC Security Groups using the ListSecurityGroups action of the AWS OIDC Integration.
func (h *Handler) awsOIDCListSecurityGroups(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCListSecurityGroupsRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
listResp, err := clt.IntegrationAWSOIDCClient().ListSecurityGroups(ctx, integrationv1.ListSecurityGroupsRequest_builder{
Integration: integrationName,
Region: req.Region,
VpcId: req.VPCID,
NextToken: req.NextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
sgs := make([]awsoidc.SecurityGroup, 0, len(listResp.GetSecurityGroups()))
for _, sg := range listResp.GetSecurityGroups() {
sgs = append(sgs, awsoidc.SecurityGroup{
Name: sg.GetName(),
ID: sg.GetId(),
Description: sg.GetDescription(),
InboundRules: awsOIDCSecurityGroupsRulesConverter(sg.GetInboundRules()),
OutboundRules: awsOIDCSecurityGroupsRulesConverter(sg.GetOutboundRules()),
})
}
return ui.AWSOIDCListSecurityGroupsResponse{
NextToken: listResp.GetNextToken(),
SecurityGroups: sgs,
}, nil
}
func awsOIDCSecurityGroupsRulesConverter(inRules []*integrationv1.SecurityGroupRule) []awsoidc.SecurityGroupRule {
out := make([]awsoidc.SecurityGroupRule, 0, len(inRules))
for _, r := range inRules {
var cidrs []awsoidc.CIDR
if len(r.GetCidrs()) > 0 {
cidrs = make([]awsoidc.CIDR, 0, len(r.GetCidrs()))
}
for _, cidr := range r.GetCidrs() {
cidrs = append(cidrs, awsoidc.CIDR{
CIDR: cidr.GetCidr(),
Description: cidr.GetDescription(),
})
}
var groupIDs []awsoidc.GroupIDRule
if len(r.GetGroupIds()) > 0 {
groupIDs = make([]awsoidc.GroupIDRule, 0, len(r.GetGroupIds()))
}
for _, group := range r.GetGroupIds() {
groupIDs = append(groupIDs, awsoidc.GroupIDRule{
GroupId: group.GetGroupId(),
Description: group.GetDescription(),
})
}
out = append(out, awsoidc.SecurityGroupRule{
IPProtocol: r.GetIpProtocol(),
FromPort: int(r.GetFromPort()),
ToPort: int(r.GetToPort()),
CIDRs: cidrs,
Groups: groupIDs,
})
}
return out
}
// awsOIDCRequiredDatabasesVPCS returns a map of required VPC's and its subnets.
// This is required during the web UI discover flow (where users opt for auto
// discovery) to determine if user can skip the auto deployment screen (where we deploy
// database agents).
//
// This api will return empty if we already have agents that can proxy the discovered databases.
// Otherwise it will return with a map of VPC and its subnets where it's values are later used
// to configure and deploy an agent (deploy an agent per unique VPC).
func (h *Handler) awsOIDCRequiredDatabasesVPCS(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCRequiredVPCSRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
respAllDatabases, err := awsOIDCListAllDatabases(ctx, clt, integrationName, req.Region)
if err != nil {
return nil, trace.Wrap(err)
}
if len(respAllDatabases) == 0 {
return nil, trace.BadParameter("there are no available RDS instances or clusters found in region %q", req.Region)
}
resp, err := awsOIDCRequiredVPCSHelper(ctx, clt, req, respAllDatabases)
if err != nil {
return nil, trace.Wrap(err)
}
return resp, nil
}
func awsOIDCListAllDatabases(ctx context.Context, clt authclient.ClientI, integration, region string) ([]*types.DatabaseV3, error) {
nextToken := ""
var fetchedRDSs []*types.DatabaseV3
// Get all rds instances.
for {
resp, err := clt.IntegrationAWSOIDCClient().ListDatabases(ctx, integrationv1.ListDatabasesRequest_builder{
Integration: integration,
Region: region,
RdsType: services.RDSDescribeTypeInstance,
Engines: []string{services.RDSEngineMySQL, services.RDSEngineMariaDB, services.RDSEnginePostgres},
NextToken: nextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
fetchedRDSs = append(fetchedRDSs, resp.GetDatabases()...)
nextToken = resp.GetNextToken()
if len(nextToken) == 0 {
break
}
}
// Get all rds clusters.
nextToken = ""
for {
resp, err := clt.IntegrationAWSOIDCClient().ListDatabases(ctx, integrationv1.ListDatabasesRequest_builder{
Integration: integration,
Region: region,
RdsType: services.RDSDescribeTypeCluster,
Engines: []string{services.RDSEngineAuroraMySQL, services.RDSEngineAuroraPostgres},
NextToken: nextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
fetchedRDSs = append(fetchedRDSs, resp.GetDatabases()...)
nextToken = resp.GetNextToken()
if len(nextToken) == 0 {
break
}
}
return fetchedRDSs, nil
}
func awsOIDCRequiredVPCSHelper(ctx context.Context, clt client.GetResourcesClient, req ui.AWSOIDCRequiredVPCSRequest, fetchedRDSs []*types.DatabaseV3) (*ui.AWSOIDCRequiredVPCSResponse, error) {
// Get all database services with ecs/fargate metadata label.
fetchedDbSvcs, err := fetchAWSOIDCDatabaseServices(ctx, clt)
if err != nil {
return nil, trace.Wrap(err)
}
// Construct map of VPCs and its subnets.
vpcLookup := map[string][]string{}
for _, db := range fetchedRDSs {
rds := db.GetAWS().RDS
vpcId := rds.VPCID
if _, found := vpcLookup[vpcId]; !found {
vpcLookup[vpcId] = rds.Subnets
continue
}
combinedSubnets := append(vpcLookup[vpcId], rds.Subnets...)
vpcLookup[vpcId] = utils.Deduplicate(combinedSubnets)
}
for _, svc := range fetchedDbSvcs {
vpcID := getDBServiceVPC(svc, req.AccountID, req.Region)
if vpcID != "" {
delete(vpcLookup, vpcID)
}
}
return &ui.AWSOIDCRequiredVPCSResponse{
VPCMapOfSubnets: vpcLookup,
}, nil
}
// awsOIDCCreateAWSAppAccess creates an AppServer that uses an AWS OIDC Integration for proxying access.
// v2 endpoint introduces "labels" field
func (h *Handler) awsOIDCCreateAWSAppAccess(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCCreateAWSAppAccessRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ig, err := clt.GetIntegration(ctx, integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
if ig.GetSubKind() != types.IntegrationSubKindAWSOIDC {
return nil, trace.BadParameter("only aws oidc integrations are supported")
}
getUserGroupLookup := h.getUserGroupLookup(r.Context(), clt)
// Reject mixed-case integration names; do not silently lowercase.
// The backend lookup is case-sensitive, so a stored record from
// before ValidIntegrationName enforced lowercase would be missed.
if strings.ToLower(integrationName) != integrationName {
return nil, trace.BadParameter("integration name %q contains uppercase characters which are no longer supported; recreate the integration with a lowercase name", integrationName)
}
publicAddr := libutils.DefaultAppPublicAddr(integrationName, h.PublicProxyAddr())
parsedRoleARN, err := awsutils.ParseRoleARN(ig.GetAWSOIDCIntegrationSpec().RoleARN)
if err != nil {
return nil, trace.Wrap(err)
}
labels := make(map[string]string)
if len(req.Labels) > 0 {
labels = req.Labels
}
labels[constants.AWSAccountIDLabel] = parsedRoleARN.AccountID
appServer, err := types.NewAppServerForAWSOIDCIntegration(integrationName, h.cfg.HostUUID, publicAddr, labels)
if err != nil {
return nil, trace.Wrap(err)
}
// If the integration name contains a dot, then the proxy must provide a certificate allowing *.<something>.<proxyPublicAddr>
if strings.Contains(integrationName, ".") {
// Teleport Cloud only provides certificates for *.<tenant>.teleport.sh, so this would generate an invalid address.
if h.GetClusterFeatures().Cloud {
return nil, trace.BadParameter(`Invalid integration name (%q) for enabling AWS Access. Please re-create the integration without the "."`, integrationName)
}
// Typically, self-hosted clusters will also have a single wildcard for the name.
// Logging a warning message should help debug the problem in case the certificate is not valid.
h.logger.WarnContext(ctx, `Enabling AWS Access using an integration with a "." might not work unless your Proxy's certificate is valid for the address`, "public_addr", appServer.GetApp().GetPublicAddr())
}
if _, err := clt.UpsertApplicationServer(ctx, appServer); err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
allowedAWSRoles, err := accessChecker.GetAllowedLoginsForResource(appServer.GetApp())
if err != nil {
return nil, trace.Wrap(err)
}
roleSet := set.New(allowedAWSRoles...)
return ui.MakeApp(appServer.GetApp(), ui.MakeAppsConfig{
LocalClusterName: h.auth.clusterName,
LocalProxyDNSName: h.proxyDNSName(),
AppClusterName: cluster.GetName(),
AWSRoles: &ui.PrincipalSet{All: roleSet, Granted: roleSet},
UserGroupLookup: getUserGroupLookup(),
Logger: h.logger,
}), nil
}
// awsOIDCDeleteAWSAppAccess deletes the AWS AppServer created that uses the AWS OIDC Integration for proxying requests.
func (h *Handler) awsOIDCDeleteAWSAppAccess(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
subkind := p.ByName("name_or_subkind")
if subkind != types.IntegrationSubKindAWSOIDC {
return nil, trace.BadParameter("only aws oidc integrations are supported")
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ig, err := clt.GetIntegration(ctx, integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
if ig.GetSubKind() != types.IntegrationSubKindAWSOIDC {
return nil, trace.BadParameter("only aws oidc integrations are supported")
}
integrationAppServer, err := h.getAppServerByName(ctx, clt, integrationName)
if err != nil {
return nil, trace.Wrap(err)
}
if integrationAppServer.GetApp().GetIntegration() != integrationName {
return nil, trace.NotFound("app %s is not using integration %s", integrationAppServer.GetName(), integrationName)
}
if err := clt.DeleteAppServer(ctx, presencev1.DeleteAppServerRequest_builder{
HostId: integrationAppServer.GetHostID(),
Name: integrationName,
Scope: integrationAppServer.GetScope(),
}.Build()); err != nil {
return nil, trace.Wrap(err)
}
return nil, nil
}
func (h *Handler) getAppServerByName(ctx context.Context, userClient authclient.ClientI, appServerName string) (types.AppServer, error) {
appServers, err := userClient.GetApplicationServers(ctx, apidefaults.Namespace)
if err != nil {
return nil, trace.Wrap(err)
}
for _, s := range appServers {
if s.GetName() == appServerName {
return s, nil
}
}
return nil, trace.NotFound("app %q not found", appServerName)
}
// awsOIDCConfigureIdP returns a script that configures AWS OIDC Integration
// by creating an OIDC Identity Provider that trusts Teleport instance.
func (h *Handler) awsOIDCConfigureIdP(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
ctx := r.Context()
queryParams := r.URL.Query()
clusterName, err := h.GetProxyClient().GetDomainName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
integrationName := queryParams.Get("integrationName")
if len(integrationName) == 0 {
return nil, trace.BadParameter("missing integrationName param")
}
// Ensure the IntegrationName is valid.
_, err = h.GetProxyClient().GetIntegration(ctx, integrationName)
// NotFound error is ignored to prevent disclosure of whether the integration exists in a public/no-auth endpoint.
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
proxyAddr, err := oidc.IssuerFromPublicAddress(h.cfg.PublicProxyAddr, "")
if err != nil {
return nil, trace.Wrap(err)
}
// The script must execute the following command:
// teleport integration configure awsoidc-idp
argsList := []string{
"integration", "configure", "awsoidc-idp",
fmt.Sprintf("--cluster=%s", shsprintf.EscapeDefaultContext(clusterName)),
fmt.Sprintf("--name=%s", shsprintf.EscapeDefaultContext(integrationName)),
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--proxy-public-url=%s", shsprintf.EscapeDefaultContext(proxyAddr)),
}
policyPreset := queryParams.Get("policyPreset")
if err := awsoidc.ValidatePolicyPreset(awsoidc.PolicyPreset(policyPreset)); err != nil {
return nil, trace.Wrap(err)
}
if policyPreset != "" {
argsList = append(argsList, fmt.Sprintf("--policy-preset=%s", shsprintf.EscapeDefaultContext(policyPreset)))
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to use the integration with AWS.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// awsOIDCConfigureListDatabasesIAM returns a script that configures the required IAM permissions to allow Listing RDS DB Clusters and Instances.
func (h *Handler) awsOIDCConfigureListDatabasesIAM(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
queryParams := r.URL.Query()
awsRegion := queryParams.Get("awsRegion")
if err := aws.IsValidRegion(awsRegion); err != nil {
return nil, trace.BadParameter("invalid awsRegion")
}
awsAccountID := queryParams.Get("awsAccountID")
if err := aws.IsValidAccountID(awsAccountID); err != nil {
return nil, trace.Wrap(err, "invalid awsAccountID")
}
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
// The script must execute the following command:
// teleport integration configure listdatabases-iam
argsList := []string{
"integration", "configure", "listdatabases-iam",
fmt.Sprintf("--aws-region=%s", shsprintf.EscapeDefaultContext(awsRegion)),
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--aws-account-id=%s", shsprintf.EscapeDefaultContext(awsAccountID)),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to complete the Database enrollment.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// accessGraphCloudSyncOIDC returns a script that configures the required IAM permissions to sync
// Cloud resources with Teleport Access Graph.
func (h *Handler) accessGraphCloudSyncOIDC(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
queryParams := r.URL.Query()
switch kind := queryParams.Get("kind"); kind {
case "aws-iam":
return h.awsAccessGraphOIDCSync(w, r, p)
default:
return nil, trace.BadParameter("unsupported kind provided %q", kind)
}
}
func (h *Handler) awsBedrockSummarizerOIDC(w http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
queryParams := r.URL.Query()
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
awsAccountID := queryParams.Get("awsAccountID")
// The script must execute the following command:
// "teleport integration configure session-summaries bedrock"
argsList := []string{
"integration", "configure", "session-summaries", "bedrock",
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--resource=%s", shsprintf.EscapeDefaultContext(queryParams.Get("resource"))),
}
if awsAccountID != "" {
argsList = append(argsList, fmt.Sprintf("--aws-account-id=%s", shsprintf.EscapeDefaultContext(awsAccountID)))
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to complete the Access Graph AWS Sync enrollment.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
func (h *Handler) awsAccessGraphOIDCSync(w http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
queryParams := r.URL.Query()
role := queryParams.Get("role")
if err := aws.IsValidIAMRoleName(role); err != nil {
return nil, trace.BadParameter("invalid role %q", role)
}
awsAccountID := queryParams.Get("awsAccountID")
if err := aws.IsValidAccountID(awsAccountID); err != nil {
return nil, trace.Wrap(err, "invalid awsAccountID")
}
// The script must execute the following command:
// "teleport integration configure access-graph aws-iam"
argsList := []string{
"integration", "configure", "access-graph", "aws-iam",
fmt.Sprintf("--role=%s", shsprintf.EscapeDefaultContext(role)),
fmt.Sprintf("--aws-account-id=%s", shsprintf.EscapeDefaultContext(awsAccountID)),
}
if sqsURL := queryParams.Get("sqsUrl"); sqsURL != "" {
if !awsoidc.IsValidSQSURL(sqsURL) {
return nil, trace.BadParameter("invalid sqsUrl %q", sqsURL)
}
argsList = append(argsList, fmt.Sprintf("--sqs-queue-url=%s", shsprintf.EscapeDefaultContext(sqsURL)))
}
if s3Bucket := queryParams.Get("cloudTrailS3Bucket"); s3Bucket != "" {
if _, err := arn.Parse(s3Bucket); err != nil {
return nil, trace.BadParameter("invalid cloudTrailS3Bucket %q", s3Bucket)
}
argsList = append(argsList, fmt.Sprintf("--cloud-trail-bucket=%s", shsprintf.EscapeDefaultContext(s3Bucket)))
}
if kmsKeysARNs := queryParams["kmsKeysARNs"]; len(kmsKeysARNs) > 0 {
for _, keyARN := range kmsKeysARNs {
if _, err := arn.Parse(keyARN); err != nil {
return nil, trace.BadParameter("invalid kmsKeysARNs %q", keyARN)
}
argsList = append(argsList, fmt.Sprintf("--kms-key=%s", shsprintf.EscapeDefaultContext(keyARN)))
}
}
if eksAuditLogs := queryParams.Get("eksAuditLogs"); eksAuditLogs != "" {
enabled, err := strconv.ParseBool(eksAuditLogs)
if err != nil {
// The error returned by ParseBool contains no more information than this
// error. As we canot wrap both it and trace.BadParameter, we do the
// latter as a preferred error type.
return nil, trace.BadParameter("invalid boolean value for eksAuditLogs %q", eksAuditLogs)
}
if enabled {
argsList = append(argsList, "--eks-audit-logs")
}
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to complete the Access Graph AWS Sync enrollment.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// awsOIDCListSubnets returns a list of VPC subnets using the ListSubnets action of the AWS OIDC Integration.
func (h *Handler) awsOIDCListSubnets(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCListSubnetsRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
listResp, err := clt.IntegrationAWSOIDCClient().ListSubnets(ctx, integrationv1.ListSubnetsRequest_builder{
Integration: integrationName,
Region: req.Region,
VpcId: req.VPCID,
NextToken: req.NextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
subnets := make([]awsoidc.Subnet, 0, len(listResp.GetSubnets()))
for _, s := range listResp.GetSubnets() {
subnets = append(subnets, awsoidc.Subnet{
Name: s.GetName(),
ID: s.GetId(),
AvailabilityZone: s.GetAvailabilityZone(),
})
}
return ui.AWSOIDCListSubnetsResponse{
NextToken: listResp.GetNextToken(),
Subnets: subnets,
}, nil
}
// awsOIDCListDatabaseVPCs returns a list of VPCs using the ListVpcs action
// of the AWS OIDC Integration, and includes a link to the ECS service if
// a database service has been deployed for each VPC.
func (h *Handler) awsOIDCListDatabaseVPCs(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req ui.AWSOIDCListVPCsRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
listResp, err := clt.IntegrationAWSOIDCClient().ListVPCs(ctx, integrationv1.ListVPCsRequest_builder{
Integration: integrationName,
Region: req.Region,
NextToken: req.NextToken,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
dbServices, err := fetchAWSOIDCDatabaseServices(ctx, clt)
if err != nil {
return nil, trace.Wrap(err)
}
serviceURLByVPC, err := getServiceURLs(dbServices, req.AccountID, req.Region, h.auth.clusterName)
if err != nil {
return nil, trace.Wrap(err)
}
vpcs := make([]ui.DatabaseEnrollmentVPC, 0, len(listResp.GetVpcs()))
for _, vpc := range listResp.GetVpcs() {
vpcs = append(vpcs, ui.DatabaseEnrollmentVPC{
VPC: awsoidc.VPC{
Name: vpc.GetName(),
ID: vpc.GetId(),
},
ECSServiceDashboardURL: serviceURLByVPC[vpc.GetId()],
})
}
return ui.AWSOIDCDatabaseVPCsResponse{
NextToken: listResp.GetNextToken(),
VPCs: vpcs,
}, nil
}
func fetchAWSOIDCDatabaseServices(ctx context.Context, clt client.GetResourcesClient) ([]types.DatabaseService, error) {
// Get all database services with the AWS OIDC agent metadata label.
var nextToken string
var fetchedDbSvcs []types.DatabaseService
for {
page, err := client.GetResourcePage[types.DatabaseService](ctx, clt, &proto.ListResourcesRequest{
ResourceType: types.KindDatabaseService,
Limit: defaults.MaxIterationLimit,
StartKey: nextToken,
Labels: map[string]string{types.AWSOIDCAgentLabel: types.True},
})
if err != nil {
return nil, trace.Wrap(err)
}
fetchedDbSvcs = append(fetchedDbSvcs, page.Resources...)
nextToken = page.NextKey
if len(nextToken) == 0 {
return fetchedDbSvcs, nil
}
}
}
// getDBServiceVPC returns the database service's VPC ID selector value if the
// database service was deployed by the AWS OIDC integration, otherwise it
// returns an empty string.
func getDBServiceVPC(svc types.DatabaseService, accountID, region string) string {
if len(svc.GetResourceMatchers()) != 1 || svc.GetResourceMatchers()[0].Labels == nil {
return ""
}
// Database services deployed by Teleport have known configurations where
// we will only define a single resource matcher.
labelMatcher := *svc.GetResourceMatchers()[0].Labels
// We check for length 3, because we are only
// wanting/checking for 3 discovery labels.
if len(labelMatcher) != 3 {
return ""
}
if slices.Compare(labelMatcher[types.DiscoveryLabelAccountID], []string{accountID}) != 0 {
return ""
}
if slices.Compare(labelMatcher[types.DiscoveryLabelRegion], []string{region}) != 0 {
return ""
}
if len(labelMatcher[types.DiscoveryLabelVPCID]) != 1 {
return ""
}
return labelMatcher[types.DiscoveryLabelVPCID][0]
}
// getServiceURLs returns a map vpcID -> service URL for ECS services deployed
// by the OIDC integration in the given account and region.
func getServiceURLs(dbServices []types.DatabaseService, accountID, region, teleportClusterName string) (map[string]string, error) {
serviceURLByVPC := make(map[string]string)
for _, svc := range dbServices {
vpcID := getDBServiceVPC(svc, accountID, region)
if vpcID == "" {
continue
}
svcURL, err := awsoidc.ECSDatabaseServiceDashboardURL(region, teleportClusterName, vpcID)
if err != nil {
return nil, trace.Wrap(err)
}
serviceURLByVPC[vpcID] = svcURL
}
return serviceURLByVPC, nil
}
// awsOIDCPing performs an health check for the integration.
// If ARN is present in the request body, that's the ARN that will be used instead of using the one stored in the integration.
// Returns meta information: account id and assumed the ARN for the IAM Role.
func (h *Handler) awsOIDCPing(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
var req ui.AWSOIDCPingRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
if req.RoleARN != "" {
integrationName = ""
}
pingResp, err := clt.IntegrationAWSOIDCClient().Ping(ctx, integrationv1.PingRequest_builder{
Integration: integrationName,
RoleArn: req.RoleARN,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSOIDCPingResponse{
AccountID: pingResp.GetAccountId(),
ARN: pingResp.GetArn(),
UserID: pingResp.GetUserId(),
}, nil
}
/*
* Teleport
* Copyright (C) 2025 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 (
"bytes"
"encoding/base64"
"fmt"
"net/http"
"strings"
"github.com/google/safetext/shsprintf"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
integrationv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/integration/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/aws"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/integrations/awscommon"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/scripts/oneoff"
"github.com/gravitational/teleport/lib/web/ui"
)
// awsRolesAnywhereConfigureTrustAnchor returns a script that configures AWS IAM Roles Anywhere Integration
// by creating:
// - IAM Roles Anywhere Trust Anchor which trusts the Teleport AWS RA CA
// - Roles Anywhere to Apps sync process:
// - IAM Role which can be assumed by the Trust Anchor and allows the APIs required by the sync process
// - IAM Roles Anywhere Profile which allows access to the IAM Role above
//
// It requires the following query parameters:
// - integrationName: the name of the AWS IAM Roles Anywhere Integration
// - trustAnchor: the name of the Trust Anchor to be created
// - syncRole: the name of the IAM Role to be created
// - syncProfile: the name of the IAM Roles Anywhere Profile to be created
func (h *Handler) awsRolesAnywhereConfigureTrustAnchor(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
ctx := r.Context()
queryParams := r.URL.Query()
integrationName := queryParams.Get("integrationName")
if integrationName == "" {
return nil, trace.BadParameter("missing integrationName param")
}
trustAnchorName := queryParams.Get("trustAnchor")
if trustAnchorName == "" {
return nil, trace.BadParameter("missing trustAnchor param")
}
if err := aws.IsValidIAMRolesAnywhereTrustAnchorName(trustAnchorName); err != nil {
return nil, trace.BadParameter("invalid trustAnchor %q", trustAnchorName)
}
syncRoleName := queryParams.Get("syncRole")
if syncRoleName == "" {
return nil, trace.BadParameter("missing syncRole param")
}
if err := aws.IsValidIAMRoleName(syncRoleName); err != nil {
return nil, trace.BadParameter("invalid role %q", syncRoleName)
}
syncProfileName := queryParams.Get("syncProfile")
if syncProfileName == "" {
return nil, trace.BadParameter("missing syncProfile param")
}
if err := aws.IsValidIAMRolesAnywhereProfileName(syncProfileName); err != nil {
return nil, trace.BadParameter("invalid syncProfile %q", syncProfileName)
}
clusterName, err := h.GetProxyClient().GetDomainName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// Ensure the IntegrationName is valid.
_, err = h.GetProxyClient().GetIntegration(ctx, integrationName)
// NotFound error is ignored to prevent disclosure of whether the integration exists in a public/no-auth endpoint.
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
authorities, err := client.ExportAllAuthorities(
ctx,
h.GetProxyClient(),
client.ExportAuthoritiesRequest{
AuthType: string(types.AWSRACA),
},
)
if err != nil {
return nil, trace.Wrap(err)
}
if len(authorities) == 0 {
return nil, trace.NotFound("no AWS IAM Roles Anywhere CA found")
}
var certAuthoritiesData [][]byte
for _, authority := range authorities {
certAuthoritiesData = append(certAuthoritiesData, authority.Data)
}
awsRACACertB64 := base64.RawStdEncoding.EncodeToString(bytes.Join(certAuthoritiesData, []byte("\n")))
// The script must execute the following command:
// teleport integration configure awsra-trust-anchor
argsList := []string{
"integration", "configure", "awsra-trust-anchor",
fmt.Sprintf("--cluster=%s", shsprintf.EscapeDefaultContext(clusterName)),
fmt.Sprintf("--name=%s", shsprintf.EscapeDefaultContext(integrationName)),
fmt.Sprintf("--trust-anchor=%s", shsprintf.EscapeDefaultContext(trustAnchorName)),
fmt.Sprintf("--sync-profile=%s", shsprintf.EscapeDefaultContext(syncProfileName)),
fmt.Sprintf("--sync-role=%s", shsprintf.EscapeDefaultContext(syncRoleName)),
fmt.Sprintf("--trust-anchor-cert-b64=%s", awsRACACertB64),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to continue the setup.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
// validateAWSRolesAnywhereIntegration performs a validation for the AWS Roles Anywhere Integration name.
// This ensures the integration name is not yet being used and that it is a valid name.
func (h *Handler) validateAWSRolesAnywhereIntegration(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
// validate integration name.
if err := awscommon.ValidIntegrationName(integrationName); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
_, err = clt.GetIntegration(ctx, integrationName)
switch {
case err == nil:
return nil, trace.AlreadyExists("integration named %q already exists", integrationName)
case trace.IsNotFound(err):
default:
return nil, trace.Wrap(err)
}
return OK(), nil
}
// awsRolesAnywherePing performs an health check for the integration.
// It returns the caller identity and the number of AWS Roles Anywhere Profiles that are active.
// If a trust anchor is provided in the body, it will be used to check the connection ignoring the integration.
// Otherwise, the integration is used to check the connection.
func (h *Handler) awsRolesAnywherePing(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
var req ui.AWSRolesAnywherePingRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
pingRequest := &integrationv1.AWSRolesAnywherePingRequest{}
// When creating an integration, the Ping is called with an empty integration, but Trust Anchor, Profile and Role ARNs must be provided.
// This allow us to check if the integration is properly configured before creating it.
switch {
case req.TrustAnchorARN != "":
if req.SyncRoleARN == "" || req.SyncProfileARN == "" {
return nil, trace.BadParameter("sync role and sync profile ARNs must be provided when trust anchor ARN is provided")
}
pingRequest.SetCustom(integrationv1.AWSRolesAnywherePingRequestWithoutIntegration_builder{
TrustAnchorArn: req.TrustAnchorARN,
RoleArn: req.SyncRoleARN,
ProfileArn: req.SyncProfileARN,
}.Build())
default:
pingRequest.SetIntegration(integrationName)
}
pingResp, err := clt.IntegrationAWSRolesAnywhereClient().AWSRolesAnywherePing(ctx, pingRequest)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.AWSRolesAnywherePingResponse{
ProfileCount: int(pingResp.GetProfileCount()),
AccountID: pingResp.GetAccountId(),
ARN: pingResp.GetArn(),
UserID: pingResp.GetUserId(),
}, nil
}
// awsRolesAnywhereListProfiles lists profiles Roles Anywhere Profiles accessible by the integration.
func (h *Handler) awsRolesAnywhereListProfiles(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
integrationName := p.ByName("name")
if integrationName == "" {
return nil, trace.BadParameter("integration name is required")
}
var req ui.AWSRolesAnywhereListProfilesRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
allProfiles := &integrationv1.ListRolesAnywhereProfilesResponse{}
var startKey string
for {
listResp, err := clt.IntegrationAWSRolesAnywhereClient().ListRolesAnywhereProfiles(ctx, integrationv1.ListRolesAnywhereProfilesRequest_builder{
Integration: integrationName,
NextPageToken: startKey,
ProfileNameFilters: req.Filters,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
allProfiles.SetProfiles(append(allProfiles.GetProfiles(), listResp.GetProfiles()...))
startKey = listResp.GetNextPageToken()
if startKey == "" {
break
}
}
return allProfiles, nil
}
// Teleport
// Copyright (C) 2024 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 (
"fmt"
"net/http"
"strings"
"github.com/google/safetext/shsprintf"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/utils/oidc"
"github.com/gravitational/teleport/lib/web/scripts/oneoff"
)
// azureOIDCConfigureIdP returns a script that configures Azure OIDC Integration
// by creating an Enterprise Application in the Azure account
// with Teleport OIDC as a trusted credential issuer.
func (h *Handler) azureOIDCConfigure(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
ctx := r.Context()
queryParams := r.URL.Query()
oidcIssuer, err := oidc.IssuerFromPublicAddress(h.cfg.PublicProxyAddr, "")
if err != nil {
return nil, trace.Wrap(err)
}
authConnectorName := queryParams.Get("authConnectorName")
if authConnectorName == "" {
return nil, trace.BadParameter("authConnectorName must be specified")
}
// Ensure the auth connector name is valid
const withSecrets = false
_, err = h.GetProxyClient().GetSAMLConnector(ctx, authConnectorName, withSecrets)
// NotFound error is ignored to prevent disclosure of whether the integration exists in a public/no-auth endpoint.
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
// The script must execute the following command:
argsList := []string{
"integration", "configure", "azure-oidc",
fmt.Sprintf("--proxy-public-addr=%s", shsprintf.EscapeDefaultContext(oidcIssuer)),
fmt.Sprintf("--auth-connector-name=%s", shsprintf.EscapeDefaultContext(authConnectorName)),
}
if tagParam := queryParams.Get("accessGraph"); tagParam != "" {
argsList = append(argsList, "--access-graph")
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to use the integration with Azure.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
/*
* Teleport
* Copyright (C) 2024 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 (
"fmt"
"net/http"
"strings"
"github.com/google/safetext/shsprintf"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/integrations/samlidp/samlidpconfig"
"github.com/gravitational/teleport/lib/web/scripts/oneoff"
)
func (h *Handler) gcpWorkforceConfigScript(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
queryParams := r.URL.Query()
samlIdPMetadataURL := fmt.Sprintf("https://%s/enterprise/saml-idp/metadata", h.PublicProxyAddr())
// validate queryParams params
if err := (samlidpconfig.GCPWorkforceAPIParams{
OrganizationID: queryParams.Get("orgId"),
PoolName: queryParams.Get("poolName"),
PoolProviderName: queryParams.Get("poolProviderName"),
SAMLIdPMetadataURL: samlIdPMetadataURL,
}).CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
// The script must execute the following command:
// teleport integration configure samlidp gcp-workforce
argsList := []string{
"integration", "configure", "samlidp", "gcp-workforce",
fmt.Sprintf("--org-id=%s", shsprintf.EscapeDefaultContext(queryParams.Get("orgId"))),
fmt.Sprintf("--pool-name=%s", shsprintf.EscapeDefaultContext(queryParams.Get("poolName"))),
fmt.Sprintf("--pool-provider-name=%s", shsprintf.EscapeDefaultContext(queryParams.Get("poolProviderName"))),
fmt.Sprintf("--idp-metadata-url=%s", shsprintf.EscapeDefaultContext(samlIdPMetadataURL)),
}
script, err := oneoff.BuildScript(oneoff.OneOffScriptParams{
EntrypointArgs: strings.Join(argsList, " "),
SuccessMessage: "Success! You can now go back to the Teleport Web UI to complete enrolling this workforce pool to Teleport SAML Identity Provider.",
})
if err != nil {
return nil, trace.Wrap(err)
}
httplib.SetScriptHeaders(w.Header())
_, err = w.Write([]byte(script))
return nil, trace.Wrap(err)
}
/*
* Teleport
* Copyright (C) 2025 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 (
"maps"
"net/http"
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
inventoryv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/inventory/v1"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
func splitQuery(value string) []string {
values := map[string]struct{}{}
for v := range strings.SplitSeq(value, ",") {
if trimmed := strings.TrimSpace(v); trimmed != "" {
values[trimmed] = struct{}{}
}
}
return slices.Collect(maps.Keys(values))
}
// listUnifiedInstancesResponse is the response for listing unified instances
type listUnifiedInstancesResponse struct {
// Instances is the list of unified instances (both instances and bot instances)
Instances []ui.UnifiedInstance `json:"instances"`
// StartKey is the next page token
StartKey string `json:"startKey"`
}
// clusterUnifiedInstancesGet returns a paginated list of unified instances
func (h *Handler) clusterUnifiedInstancesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
startKey := values.Get("startKey")
// Default values
sort := inventoryv1.UnifiedInstanceSort_UNIFIED_INSTANCE_SORT_NAME
order := inventoryv1.SortOrder_SORT_ORDER_ASCENDING
if sortParam := values.Get("sort"); sortParam != "" {
parts := strings.SplitN(sortParam, ":", 2)
fieldName := strings.ToLower(parts[0])
switch fieldName {
case "name", "hostname":
sort = inventoryv1.UnifiedInstanceSort_UNIFIED_INSTANCE_SORT_NAME
case "type":
sort = inventoryv1.UnifiedInstanceSort_UNIFIED_INSTANCE_SORT_TYPE
case "version":
sort = inventoryv1.UnifiedInstanceSort_UNIFIED_INSTANCE_SORT_VERSION
}
if len(parts) == 2 {
direction := strings.ToLower(parts[1])
if direction == "desc" || direction == "descending" {
order = inventoryv1.SortOrder_SORT_ORDER_DESCENDING
}
}
}
// We don't use splitQuery for parsing upgraders since it should be possible to filter for upgrader "" (none),
// and splitQuery would remove it.
var upgraders []string
if upgradersParam := values.Get("upgraders"); upgradersParam != "" {
upgraders = strings.Split(upgradersParam, ",")
}
filter := inventoryv1.ListUnifiedInstancesFilter_builder{
Search: values.Get("search"),
PredicateExpression: values.Get("query"),
Services: splitQuery(values.Get("services")),
Upgraders: upgraders,
UpdaterGroups: splitQuery(values.Get("updaterGroups")),
}.Build()
var hasInstance, hasBotInstance bool
for t := range strings.SplitSeq(values.Get("types"), ",") {
switch strings.ToLower(strings.TrimSpace(t)) {
case "instance":
hasInstance = true
case "bot_instance":
hasBotInstance = true
}
if hasInstance && hasBotInstance {
break
}
}
if hasInstance {
filter.SetInstanceTypes(append(filter.GetInstanceTypes(), inventoryv1.InstanceType_INSTANCE_TYPE_INSTANCE))
}
if hasBotInstance {
filter.SetInstanceTypes(append(filter.GetInstanceTypes(), inventoryv1.InstanceType_INSTANCE_TYPE_BOT_INSTANCE))
}
resp, err := clt.ListUnifiedInstances(r.Context(), inventoryv1.ListUnifiedInstancesRequest_builder{
PageSize: limit,
PageToken: startKey,
Sort: sort,
Order: order,
Filter: filter,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
uiInstances := make([]ui.UnifiedInstance, 0, len(resp.GetItems()))
for _, item := range resp.GetItems() {
uiInstances = append(uiInstances, ui.MakeUnifiedInstance(item))
}
return &listUnifiedInstancesResponse{
Instances: uiInstances,
StartKey: resp.GetNextPageToken(),
}, nil
}
/*
* 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"
"encoding/hex"
"fmt"
"hash/fnv"
"net/http"
"net/url"
"reflect"
"sort"
"strconv"
"strings"
"time"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/scopes/joining"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/ui"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/web/scripts"
webui "github.com/gravitational/teleport/lib/web/ui"
)
const (
HeaderTokenName = "X-Teleport-TokenName"
)
// nodeJoinToken contains node token fields for the UI.
type nodeJoinToken struct {
// ID is token ID.
ID string `json:"id"`
// Expiry is token expiration time.
Expiry time.Time `json:"expiry"`
// Method is the join method that the token supports
Method types.JoinMethod `json:"method"`
// SuggestedLabels contains the set of labels we expect the node to set when using this token
SuggestedLabels []ui.Label `json:"suggestedLabels,omitempty"`
}
// scriptSettings is used to hold values which are passed into the function that
// generates the join script.
type scriptSettings struct {
token string
appInstallMode bool
appName string
appURI string
joinMethod string
databaseInstallMode bool
discoveryInstallMode bool
discoveryGroup string
}
// automaticUpgrades returns whether automaticUpgrades should be enabled.
func automaticUpgrades(features proto.Features) bool {
return features.AutomaticUpgrades && features.Cloud
}
// Currently we aren't paginating this endpoint as we don't
// expect many tokens to exist at a time. I'm leaving it in a "paginated" form
// without a nextKey for now so implementing pagination won't change the response shape
// TODO (avatus) implement pagination
// GetTokensResponse returns a list of JoinTokens.
type GetTokensResponse struct {
Items []webui.JoinToken `json:"items"`
}
func (h *Handler) getTokens(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
// This endpoint returns provision, static, and user tokens together for
// compatibility with the legacy all-in-one "GetTokens" RPC.
staticTokens, err := clt.GetStaticTokens(r.Context())
if err != nil {
return nil, trace.Wrap(err, "getting static tokens")
}
tokens := staticTokens.GetStaticTokens()
provisionTokens, err := stream.Collect(clientutils.Resources(r.Context(),
func(ctx context.Context, pageSize int, pageKey string) ([]types.ProvisionToken, string, error) {
return clt.ListProvisionTokens(ctx, pageSize, pageKey, nil, "")
},
))
if err != nil {
return nil, trace.Wrap(err, "getting provision tokens")
}
tokens = append(tokens, provisionTokens...)
userTokens, err := stream.Collect(clientutils.Resources(r.Context(), clt.ListResetPasswordTokens))
if err != nil {
return nil, trace.Wrap(err, "getting user tokens")
}
// Converting the user tokens as provision tokens for presentation and
// backward compatibility.
for _, t := range userTokens {
roles := types.SystemRoles{types.RoleSignup}
tok, err := types.NewProvisionToken(t.GetName(), roles, t.Expiry())
if err != nil {
return nil, trace.Wrap(err, "converting user token as a provision token")
}
tokens = append(tokens, tok)
}
uiTokens, err := webui.MakeJoinTokens(tokens)
if err != nil {
return nil, trace.Wrap(err)
}
return GetTokensResponse{
Items: uiTokens,
}, nil
}
// ListProvisionTokensResponse contains a paginated list of provision tokens.
type ListProvisionTokensResponse struct {
Items []webui.JoinToken `json:"items"`
NextPageToken string `json:"next_page_token,omitempty"`
}
// listProvisionTokens returns a paginated list of provision tokens. Items can
// be filtered by role and bot name. Tokens with ANY of the provided roles are
// returned. If a bot name is provided, only tokens having a role of Bot are
// returned.
func (h *Handler) listProvisionTokens(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
var pageSize int64 = 20
if r.URL.Query().Has("page_size") {
pageSize, err = strconv.ParseInt(r.URL.Query().Get("page_size"), 10, 32)
if err != nil {
return nil, trace.Wrap(err, "failed to parse page_size")
}
}
roles, err := types.NewTeleportRoles(r.URL.Query()["role"])
if err != nil {
return nil, trace.Wrap(err)
}
items, nextToken, err := clt.ListProvisionTokens(r.Context(), int(pageSize), r.URL.Query().Get("page_token"), roles, r.URL.Query().Get("bot_name"))
if err != nil {
return nil, trace.Wrap(err)
}
uiTokens, err := webui.MakeJoinTokens(items)
if err != nil {
return nil, trace.Wrap(err)
}
return ListProvisionTokensResponse{
Items: uiTokens,
NextPageToken: nextToken,
}, nil
}
func (h *Handler) deleteToken(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
token := r.Header.Get(HeaderTokenName)
if token == "" {
return nil, trace.BadParameter("requires a token to delete")
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
if err := clt.DeleteToken(r.Context(), token); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
type CreateTokenRequest struct {
Content string `json:"content"`
}
func (h *Handler) updateTokenYAML(w http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
tokenId := r.Header.Get(HeaderTokenName)
if tokenId == "" {
return nil, trace.BadParameter("requires a token name to edit")
}
var yaml CreateTokenRequest
if err := httplib.ReadResourceJSON(r, &yaml); err != nil {
return nil, trace.Wrap(err)
}
extractedRes, err := ExtractResourceAndValidate(yaml.Content)
if err != nil {
return nil, trace.Wrap(err)
}
if tokenId != extractedRes.Metadata.Name {
return nil, trace.BadParameter("renaming tokens is not supported")
}
token, err := services.UnmarshalProvisionToken(extractedRes.Raw)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
err = clt.UpsertToken(r.Context(), token)
if err != nil {
return nil, trace.Wrap(err)
}
uiToken, err := webui.MakeJoinToken(token)
if err != nil {
return nil, trace.Wrap(err)
}
return uiToken, trace.Wrap(err)
}
type upsertTokenHandleRequest struct {
types.ProvisionTokenSpecV2
Name string `json:"name"`
}
func (h *Handler) upsertTokenHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
var req upsertTokenHandleRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
var existingToken types.ProvisionToken
if req.Name != "" {
existingToken, err = clt.GetToken(r.Context(), req.Name)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
}
var expires time.Time
switch req.JoinMethod {
case types.JoinMethodGCP, types.JoinMethodIAM, types.JoinMethodOracle, types.JoinMethodGitHub, types.JoinMethodGitLab:
// IAM, GCP, Oracle, GitHub and GitLab tokens should never expire.
expires = time.Time{}
default:
// Set expires time to default node join token TTL.
expires = time.Now().UTC().Add(defaults.NodeJoinTokenTTL)
}
name := req.Name
if name == "" {
randName, err := utils.CryptoRandomHex(defaults.TokenLenBytes)
if err != nil {
return nil, trace.Wrap(err)
}
name = randName
}
token, err := types.NewProvisionTokenFromSpec(name, expires, req.ProvisionTokenSpecV2)
if err != nil {
return nil, trace.Wrap(err)
}
// If this is an edit, then overwrite the metadata to retain the existing fields
if existingToken != nil {
token.SetMetadata(existingToken.GetMetadata())
}
err = clt.UpsertToken(r.Context(), token)
if err != nil {
return nil, trace.Wrap(err)
}
uiToken, err := webui.MakeJoinToken(token)
if err != nil {
return nil, trace.Wrap(err)
}
return uiToken, nil
}
// createTokenForDiscoveryHandle creates tokens used during guided discover flows.
// V2 endpoint processes "suggestedLabels" field.
func (h *Handler) createTokenForDiscoveryHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
var req types.ProvisionTokenSpecV2
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
var expires time.Time
var tokenName string
switch req.JoinMethod {
case types.JoinMethodIAM:
// to prevent generation of redundant IAM tokens
// we generate a deterministic name for them
tokenName, err = generateIAMTokenName(req.Allow)
if err != nil {
return nil, trace.Wrap(err)
}
// if a token with this name is found and it has indeed the same rule set,
// return it. Otherwise, go ahead and create it
t, err := clt.GetToken(r.Context(), tokenName)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
if err == nil {
// check if the token found has the right rules
if t.GetJoinMethod() != types.JoinMethodIAM || !isSameRuleSet(req.Allow, t.GetAWSAllowRules()) {
return nil, trace.BadParameter("failed to create token: token with name %q already exists and does not have the expected allow rules", tokenName)
}
return &nodeJoinToken{
ID: t.GetName(),
Expiry: t.Expiry(),
Method: t.GetJoinMethod(),
}, nil
}
// IAM tokens should 'never' expire
expires = time.Now().UTC().AddDate(1000, 0, 0)
case types.JoinMethodAzure:
tokenName, err := generateAzureTokenName(req.Azure.Allow)
if err != nil {
return nil, trace.Wrap(err)
}
t, err := clt.GetToken(r.Context(), tokenName)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
v2token, ok := t.(*types.ProvisionTokenV2)
if !ok {
return nil, trace.BadParameter("Azure join requires v2 token")
}
if err == nil {
if t.GetJoinMethod() != types.JoinMethodAzure || !isSameAzureRuleSet(req.Azure.Allow, v2token.Spec.Azure.Allow) {
return nil, trace.BadParameter("failed to create token: token with name %q already exists and does not have the expected allow rules", tokenName)
}
return &nodeJoinToken{
ID: t.GetName(),
Expiry: t.Expiry(),
Method: t.GetJoinMethod(),
}, nil
}
default:
tokenName, err = utils.CryptoRandomHex(defaults.TokenLenBytes)
if err != nil {
return nil, trace.Wrap(err)
}
expires = time.Now().UTC().Add(defaults.NodeJoinTokenTTL)
}
// If using the automatic method to add a Node, the `install.sh` will add the token's suggested labels
// as part of the initial Labels configuration for that Node
// Script install-node.sh:
// ...
// $ teleport configure ... --labels <suggested_label=value>,<suggested_label=value> ...
// ...
//
// We create an ID and return it as part of the Token, so the UI can use this ID to query the Node that joined using this token
// WebUI can then query the resources by this id and answer the question:
// - Which Node joined the cluster from this token Y?
if req.SuggestedLabels == nil {
req.SuggestedLabels = make(types.Labels)
}
req.SuggestedLabels[types.InternalResourceIDLabel] = apiutils.Strings{uuid.NewString()}
provisionToken, err := types.NewProvisionTokenFromSpec(tokenName, expires, req)
if err != nil {
return nil, trace.Wrap(err)
}
err = clt.CreateToken(r.Context(), provisionToken)
if err != nil {
return nil, trace.Wrap(err)
}
suggestedLabels := make([]ui.Label, 0, len(req.SuggestedLabels))
for labelKey, labelValues := range req.SuggestedLabels {
suggestedLabels = append(suggestedLabels, ui.Label{
Name: labelKey,
Value: strings.Join(labelValues, " "),
})
}
return &nodeJoinToken{
ID: tokenName,
Expiry: expires,
Method: provisionToken.GetJoinMethod(),
SuggestedLabels: suggestedLabels,
}, nil
}
func (h *Handler) getNodeJoinScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (any, error) {
httplib.SetScriptHeaders(w.Header())
settings := scriptSettings{
token: params.ByName("token"),
appInstallMode: false,
joinMethod: r.URL.Query().Get("method"),
}
script, err := h.getJoinScript(r.Context(), settings)
if err != nil {
h.logger.InfoContext(r.Context(), "Failed to return the node install script", "error", err)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
w.WriteHeader(http.StatusOK)
if _, err := fmt.Fprintln(w, script); err != nil {
h.logger.InfoContext(r.Context(), "Failed to return the node install script", "error", err)
w.Write(scripts.ErrorBashScript)
}
return nil, nil
}
func (h *Handler) getAppJoinScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (any, error) {
httplib.SetScriptHeaders(w.Header())
queryValues := r.URL.Query()
name, err := url.QueryUnescape(queryValues.Get("name"))
if err != nil {
h.logger.DebugContext(r.Context(), "Failed to return the app install script",
"query_param", "name",
"error", err,
)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
uri, err := url.QueryUnescape(queryValues.Get("uri"))
if err != nil {
h.logger.DebugContext(r.Context(), "Failed to return the app install script",
"query_param", "uri",
"error", err,
)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
settings := scriptSettings{
token: params.ByName("token"),
appInstallMode: true,
appName: name,
appURI: uri,
}
script, err := h.getJoinScript(r.Context(), settings)
if err != nil {
h.logger.InfoContext(r.Context(), "Failed to return the app install script", "error", err)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
w.WriteHeader(http.StatusOK)
if _, err := fmt.Fprintln(w, script); err != nil {
h.logger.DebugContext(r.Context(), "Failed to return the app install script", "error", err)
w.Write(scripts.ErrorBashScript)
}
return nil, nil
}
func (h *Handler) getDatabaseJoinScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (any, error) {
httplib.SetScriptHeaders(w.Header())
settings := scriptSettings{
token: params.ByName("token"),
databaseInstallMode: true,
}
script, err := h.getJoinScript(r.Context(), settings)
if err != nil {
h.logger.InfoContext(r.Context(), "Failed to return the database install script", "error", err)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
w.WriteHeader(http.StatusOK)
if _, err := fmt.Fprintln(w, script); err != nil {
h.logger.DebugContext(r.Context(), "Failed to return the database install script", "error", err)
w.Write(scripts.ErrorBashScript)
}
return nil, nil
}
func (h *Handler) getDiscoveryJoinScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (any, error) {
httplib.SetScriptHeaders(w.Header())
queryValues := r.URL.Query()
const discoveryGroupQueryParam = "discoveryGroup"
discoveryGroup, err := url.QueryUnescape(queryValues.Get(discoveryGroupQueryParam))
if err != nil {
h.logger.DebugContext(r.Context(), "Failed to return the discovery install script",
"error", err,
"query_param", discoveryGroupQueryParam,
)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
if discoveryGroup == "" {
h.logger.DebugContext(r.Context(), "Failed to return the discovery install script. Missing required fields",
"query_param", discoveryGroupQueryParam,
)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
settings := scriptSettings{
token: params.ByName("token"),
discoveryInstallMode: true,
discoveryGroup: discoveryGroup,
}
script, err := h.getJoinScript(r.Context(), settings)
if err != nil {
h.logger.InfoContext(r.Context(), "Failed to return the discovery install script", "error", err)
w.Write(scripts.ErrorBashScript)
return nil, nil
}
w.WriteHeader(http.StatusOK)
if _, err := fmt.Fprintln(w, script); err != nil {
h.logger.DebugContext(r.Context(), "Failed to return the discovery install script", "error", err)
w.Write(scripts.ErrorBashScript)
}
return nil, nil
}
func (h *Handler) getJoinScript(ctx context.Context, settings scriptSettings) (string, error) {
joinMethod := types.JoinMethod(settings.joinMethod)
switch joinMethod {
case types.JoinMethodUnspecified, types.JoinMethodToken:
if err := validateJoinToken(settings.token); err != nil {
return "", trace.Wrap(err)
}
case types.JoinMethodIAM:
default:
return "", trace.BadParameter("join method %q is not supported via script", settings.joinMethod)
}
clt := h.GetProxyClient()
// The provided token can be attacker controlled, so we must validate
// it with the backend before using it to generate the script.
token, err := clt.GetToken(ctx, settings.token)
if err != nil {
return "", trace.BadParameter("invalid token")
}
// TODO(hugoShaka): hit the local accesspoint which has a cache instead of asking the auth every time.
// Get the CA pin hashes of the cluster to join.
localCAResponse, err := clt.GetClusterCACert(ctx)
if err != nil {
return "", trace.Wrap(err)
}
caPins, err := tlsca.CalculatePins(localCAResponse.TLSCA)
if err != nil {
return "", trace.Wrap(err)
}
installOpts, err := h.installScriptOptions(ctx)
if err != nil {
return "", trace.Wrap(err, "Building install script options")
}
tokenName := token.GetName()
if secret, ok := token.GetSecret(); ok {
tokenName = joining.EncodeScopedToken(tokenName, secret)
}
nodeInstallOpts := scripts.InstallNodeScriptOptions{
InstallOptions: installOpts,
Token: tokenName,
CAPins: caPins,
// We are using the joinMethod from the script settings instead of the one from the token
// to reproduce the previous script behavior. I'm also afraid that using the
// join method from the token would provide an oracle for an attacker wanting to discover
// the join method.
// We might want to change this in the future to lookup the join method from the token
// to avoid potential mismatch and allow the caller to not care about the join method.
JoinMethod: joinMethod,
Labels: token.GetSuggestedLabels(),
LabelMatchers: token.GetSuggestedAgentMatcherLabels(),
AppServiceEnabled: settings.appInstallMode,
AppName: settings.appName,
AppURI: settings.appURI,
DatabaseServiceEnabled: settings.databaseInstallMode,
DiscoveryServiceEnabled: settings.discoveryInstallMode,
DiscoveryGroup: settings.discoveryGroup,
}
return scripts.GetNodeInstallScript(ctx, nodeInstallOpts)
}
// validateJoinToken validate a join token.
func validateJoinToken(token string) error {
decodedToken, err := hex.DecodeString(token)
if err != nil {
return trace.BadParameter("invalid token %q", token)
}
if len(decodedToken) != defaults.TokenLenBytes {
return trace.BadParameter("invalid token %q", decodedToken)
}
return nil
}
// generateIAMTokenName makes a deterministic name for a iam join token
// based on its rule set
func generateIAMTokenName(rules []*types.TokenRule) (string, error) {
// sort the rules by (account ID, arn)
// to make sure a set of rules will produce the same hash,
// no matter the order they are in the slice
orderedRules := make([]*types.TokenRule, len(rules))
copy(orderedRules, rules)
sortRules(orderedRules)
h := fnv.New32a()
for _, r := range orderedRules {
s := fmt.Sprintf("%s%s", r.AWSAccount, r.AWSARN)
_, err := h.Write([]byte(s))
if err != nil {
return "", trace.Wrap(err)
}
}
return fmt.Sprintf("teleport-ui-iam-%d", h.Sum32()), nil
}
// generateAzureTokenName makes a deterministic name for an azure join token
// based on its rule set.
func generateAzureTokenName(rules []*types.ProvisionTokenSpecV2Azure_Rule) (string, error) {
orderedRules := make([]*types.ProvisionTokenSpecV2Azure_Rule, len(rules))
copy(orderedRules, rules)
sortAzureRules(orderedRules)
h := fnv.New32a()
for _, r := range orderedRules {
hashInput := r.Subscription
if r.Tenant != "" {
hashInput += "tenant:" + r.Tenant
}
_, err := h.Write([]byte(hashInput))
if err != nil {
return "", trace.Wrap(err)
}
}
return fmt.Sprintf("teleport-ui-azure-%d", h.Sum32()), nil
}
// sortRules sorts a slice of rules based on their AWS Account ID and ARN
func sortRules(rules []*types.TokenRule) {
sort.Slice(rules, func(i, j int) bool {
iAcct, jAcct := rules[i].AWSAccount, rules[j].AWSAccount
// if accountID is the same, sort based on arn
if iAcct == jAcct {
arn1, arn2 := rules[i].AWSARN, rules[j].AWSARN
return arn1 < arn2
}
return iAcct < jAcct
})
}
// sortAzureRules sorts a slice of Azure rules based on their subscription and tenant.
func sortAzureRules(rules []*types.ProvisionTokenSpecV2Azure_Rule) {
sort.Slice(rules, func(i, j int) bool {
if rules[i].Tenant != rules[j].Tenant {
return rules[i].Tenant < rules[j].Tenant
}
return rules[i].Subscription < rules[j].Subscription
})
}
// isSameRuleSet check if r1 and r2 are the same rules, ignoring the order
func isSameRuleSet(r1 []*types.TokenRule, r2 []*types.TokenRule) bool {
sortRules(r1)
sortRules(r2)
return reflect.DeepEqual(r1, r2)
}
// isSameAzureRuleSet checks if r1 and r2 are the same rules, ignoring order.
func isSameAzureRuleSet(r1, r2 []*types.ProvisionTokenSpecV2Azure_Rule) bool {
sortAzureRules(r1)
sortAzureRules(r2)
return reflect.DeepEqual(r1, r2)
}
// Teleport
// Copyright (C) 2024 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"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/jwt"
)
func (h *Handler) jwks(ctx context.Context, caType types.CertAuthType, includeBlankKeyID bool) (*JWKSResponse, error) {
clusterName, err := h.GetProxyClient().GetDomainName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
// Fetch the JWT public keys only.
ca, err := h.GetProxyClient().GetCertAuthority(ctx, types.CertAuthID{
Type: caType,
DomainName: clusterName,
}, false /* loadKeys */)
if err != nil {
return nil, trace.Wrap(err)
}
pairs := ca.GetTrustedJWTKeyPairs()
// Create response and allocate space for the keys.
var resp JWKSResponse
resp.Keys = make([]jwt.JWK, 0, len(pairs))
// Loop over and all add public keys in JWK format.
for _, key := range pairs {
jwk, err := jwt.MarshalJWK(key.PublicKey)
if err != nil {
return nil, trace.Wrap(err)
}
resp.Keys = append(resp.Keys, jwk)
// Return an additional copy of the same JWK
// with KeyID set to the empty string for compatibility.
if includeBlankKeyID {
jwk.KeyID = ""
resp.Keys = append(resp.Keys, jwk)
}
}
return &resp, nil
}
/*
* Teleport
* Copyright (C) 2024 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 (
"bytes"
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"log/slog"
"net/http"
"strings"
"sync/atomic"
"time"
"github.com/gogo/protobuf/proto"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
oteltrace "go.opentelemetry.io/otel/trace"
v1 "k8s.io/api/core/v1"
"k8s.io/client-go/kubernetes"
"k8s.io/client-go/kubernetes/scheme"
"k8s.io/client-go/rest"
"k8s.io/client-go/tools/remotecommand"
clientproto "github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/session"
logutils "github.com/gravitational/teleport/lib/utils/log"
"github.com/gravitational/teleport/lib/web/terminal"
)
// podExecHandler connects Kube exec session and web-based terminal via a websocket.
type podExecHandler struct {
teleportCluster string
configTLSServerName string
configServerAddr string
publicProxyAddr string
req *PodExecRequest
sess session.Session
sctx *SessionContext
ws *websocket.Conn
keepAliveInterval time.Duration
logger *slog.Logger
userClient authclient.ClientI
localCA types.CertAuthority
// closedByClient indicates if the websocket connection was closed by the
// user (closing the browser tab, exiting the session, etc).
closedByClient atomic.Bool
}
// PodExecRequest describes a request to create a web-based terminal
// to exec into a pod.
type PodExecRequest struct {
// KubeCluster specifies what Kubernetes cluster to connect to.
KubeCluster string `json:"kubeCluster"`
// Namespace is the namespace of the target pod
Namespace string `json:"namespace"`
// Pod is the target pod to connect to.
Pod string `json:"pod"`
// Container is a container within the target pod to connect to, optional.
Container string `json:"container"`
// Command is the command to run at the target pod.
Command string `json:"command"`
// IsInteractive specifies whether exec request should have interactive TTY.
IsInteractive bool `json:"isInteractive"`
// Term is the initial PTY size.
Term session.TerminalParams `json:"term"`
}
func (r *PodExecRequest) Validate() error {
if r.KubeCluster == "" {
return trace.BadParameter("missing parameter KubeCluster")
}
if r.Namespace == "" {
return trace.BadParameter("missing parameter Namespace")
}
if r.Pod == "" {
return trace.BadParameter("missing parameter Pod")
}
if r.Command == "" {
return trace.BadParameter("missing parameter Command")
}
if len(r.Namespace) > 63 {
return trace.BadParameter("Namespace is too long, maximum length is 63 characters")
}
if len(r.Pod) > 63 {
return trace.BadParameter("Pod is too long, maximum length is 63 characters")
}
if len(r.Container) > 63 {
return trace.BadParameter("Container is too long, maximum length is 63 characters")
}
if len(r.Command) > 10000 {
return trace.BadParameter("Command is too long, maximum length is 10000 characters")
}
return nil
}
// ServeHTTP sends session metadata to web UI to signal beginning of the session, then
// handles Kube exec request and connects it to web based terminal input/output.
func (p *podExecHandler) ServeHTTP(_ http.ResponseWriter, r *http.Request) {
// Allow closing websocket if the user logs out before exiting
// the session.
p.sctx.AddClosers(p)
defer p.sctx.RemoveCloser(p)
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: p.sess})
if err != nil {
p.logger.ErrorContext(r.Context(), "failed marshaling session data", "error", err)
if err := p.sendErrorMessage(err); err != nil {
p.logger.ErrorContext(r.Context(), "failed to send error message to client", "error", err)
}
return
}
envelope := &terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketSessionMetadata,
Payload: string(sessionMetadataResponse),
}
envelopeBytes, err := proto.Marshal(envelope)
if err != nil {
p.logger.ErrorContext(r.Context(), "failed marshaling message envelope", "error", err)
if err := p.sendErrorMessage(err); err != nil {
p.logger.ErrorContext(r.Context(), "failed to send error message to client", "error", err)
}
return
}
err = p.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes)
if err != nil {
p.logger.ErrorContext(r.Context(), "failed write session data message", "error", err)
if err := p.sendErrorMessage(err); err != nil {
p.logger.ErrorContext(r.Context(), "failed to send error message to client", "error", err)
}
return
}
if err := p.handler(r); err != nil {
p.logger.ErrorContext(r.Context(), "handling kube session unexpectedly terminated", "error", err)
if err := p.sendErrorMessage(err); err != nil {
p.logger.ErrorContext(r.Context(), "failed to send error message to client", "error", err)
}
}
}
func (p *podExecHandler) Close() error {
return trace.Wrap(p.ws.Close())
}
func (p *podExecHandler) sendErrorMessage(err error) error {
if p.closedByClient.Load() {
return nil
}
envelope := &terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketError,
Payload: err.Error(),
}
envelopeBytes, err := proto.Marshal(envelope)
if err != nil {
return trace.Wrap(err, "creating envelope payload")
}
if err := p.ws.WriteMessage(websocket.BinaryMessage, envelopeBytes); err != nil {
return trace.Wrap(err, "writing error message")
}
return nil
}
func (p *podExecHandler) handler(r *http.Request) error {
p.logger.DebugContext(r.Context(), "Creating websocket stream for a kube exec request")
// Create a context for signaling when the terminal session is over and
// link it first with the trace context from the request context
tctx := oteltrace.ContextWithRemoteSpanContext(context.Background(), oteltrace.SpanContextFromContext(r.Context()))
ctx, cancel := context.WithCancel(tctx)
defer cancel()
defaultCloseHandler := p.ws.CloseHandler()
p.ws.SetCloseHandler(func(code int, text string) error {
p.closedByClient.Store(true)
p.logger.DebugContext(r.Context(), "websocket connection was closed by client")
cancel()
// Call the default close handler if one was set.
if defaultCloseHandler != nil {
err := defaultCloseHandler(code, text)
return trace.NewAggregate(err, p.Close())
}
return trace.Wrap(p.Close())
})
// Start sending ping frames through websocket to the client.
go startWSPingLoop(r.Context(), p.ws, p.keepAliveInterval, p.logger, p.Close)
pk, err := keys.ParsePrivateKey(p.sctx.cfg.Session.GetTLSPriv())
if err != nil {
return trace.Wrap(err, "failed getting user private key from the session")
}
privateKeyPEM, err := pk.SoftwarePrivateKeyPEM()
if err != nil {
return trace.Wrap(err, "failed getting software private key")
}
publicKeyPEM, err := keys.MarshalPublicKey(pk.Public())
if err != nil {
return trace.Wrap(err, "failed to marshal public key")
}
resizeQueue := newTermSizeQueue(ctx, remotecommand.TerminalSize{
Width: p.req.Term.Winsize().Width,
Height: p.req.Term.Winsize().Height,
})
stream := terminal.NewStream(ctx, terminal.StreamConfig{
WS: p.ws,
Logger: p.logger,
Handlers: map[string]terminal.WSHandlerFunc{
defaults.WebsocketResize: p.handleResize(resizeQueue),
},
})
certsReq := clientproto.UserCertsRequest{
TLSPublicKey: publicKeyPEM,
Username: p.sctx.GetUser(),
Expires: p.sctx.cfg.Session.GetExpiryTime(),
Format: constants.CertificateFormatStandard,
RouteToCluster: p.teleportCluster,
KubernetesCluster: p.req.KubeCluster,
Usage: clientproto.UserCertsRequest_Kubernetes,
}
var certs *clientproto.Certs
result, err := client.PerformSessionMFACeremony(ctx, client.PerformSessionMFACeremonyParams{
CurrentAuthClient: p.userClient,
RootAuthClient: p.sctx.cfg.RootClient,
MFACeremony: newMFACeremony(stream.WSStream, p.sctx.cfg.RootClient.CreateAuthenticateChallenge, p.publicProxyAddr),
MFAAgainstRoot: p.sctx.cfg.RootClusterName == p.teleportCluster,
MFARequiredReq: &clientproto.IsMFARequiredRequest{
Target: &clientproto.IsMFARequiredRequest_KubernetesCluster{KubernetesCluster: p.req.KubeCluster},
},
CertsReq: &certsReq,
})
if err != nil && !errors.Is(err, services.ErrSessionMFANotRequired) {
return trace.Wrap(err, "failed performing mfa ceremony")
} else if result != nil {
certs = result.NewCerts
}
if certs == nil {
certs, err = p.sctx.cfg.RootClient.GenerateUserCerts(ctx, certsReq)
if err != nil {
return trace.Wrap(err, "failed issuing user certs")
}
}
restConfig, err := createKubeRestConfig(p.configServerAddr, p.configTLSServerName, p.localCA, certs.TLS, privateKeyPEM)
if err != nil {
return trace.Wrap(err, "failed creating Kubernetes rest config")
}
kubeClient, err := kubernetes.NewForConfig(restConfig)
if err != nil {
return trace.Wrap(err, "failed creating Kubernetes client")
}
kubeReq := kubeClient.CoreV1().RESTClient().Post().Resource("pods").Name(p.req.Pod).
Namespace(p.req.Namespace).SubResource("exec")
option := &v1.PodExecOptions{
Container: p.req.Container,
Command: strings.Split(p.req.Command, " "),
TTY: p.req.IsInteractive,
Stdin: p.req.IsInteractive,
Stdout: true,
Stderr: !p.req.IsInteractive,
}
kubeReq.VersionedParams(option, scheme.ParameterCodec)
p.logger.DebugContext(ctx, "Web kube exec request created", "url", logutils.StringerAttr(kubeReq.URL()))
wsExec, err := remotecommand.NewWebSocketExecutor(restConfig, "POST", kubeReq.URL().String())
if err != nil {
return trace.Wrap(err, "failed creating websocket executor")
}
streamOpts := remotecommand.StreamOptions{
Stdin: stream,
Stdout: stream,
Tty: p.req.IsInteractive,
TerminalSizeQueue: resizeQueue,
}
if !p.req.IsInteractive {
streamOpts.Stderr = stderrWriter{stream: stream}
}
if err := wsExec.StreamWithContext(ctx, streamOpts); err != nil {
return trace.Wrap(err, "failed exec command streaming")
}
if p.closedByClient.Load() {
return nil // No need to send close envelope to the web UI if it was already closed by user.
}
// TODO(anton): refactor UI - right now if we send the close message UI will remove all text
// from the document, which doesn't make sense for non-interactive command, since user
// never has the chance to see the output.
if p.req.IsInteractive {
// Send close envelope to web terminal upon exit without an error.
if err := stream.SendCloseMessage(""); err != nil {
p.logger.ErrorContext(ctx, "unable to send close event to web client", "error", err)
}
}
if err := stream.Close(); err != nil {
p.logger.ErrorContext(ctx, "unable to close websocket stream to web client", "error", err)
return nil
}
p.logger.DebugContext(ctx, "Sent close event to web client", "error", err)
return nil
}
func (p *podExecHandler) handleResize(termSizeQueue *termSizeQueue) func(context.Context, terminal.Envelope) {
return func(ctx context.Context, envelope terminal.Envelope) {
var e map[string]any
if err := json.Unmarshal([]byte(envelope.Payload), &e); err != nil {
p.logger.WarnContext(ctx, "Failed to parse resize payload", "error", err)
return
}
size, ok := e["size"].(string)
if !ok {
p.logger.ErrorContext(ctx, "got unexpected size type, expected type string", "size_type", logutils.TypeAttr(size))
return
}
params, err := session.UnmarshalTerminalParams(size)
if err != nil {
p.logger.WarnContext(ctx, "Failed to retrieve terminal size", "error", err)
return
}
// nil params indicates the channel was closed
if params == nil {
return
}
termSizeQueue.AddSize(remotecommand.TerminalSize{
Width: params.Winsize().Width,
Height: params.Winsize().Height,
})
}
}
type termSizeQueue struct {
incoming chan remotecommand.TerminalSize
ctx context.Context
}
func newTermSizeQueue(ctx context.Context, initialSize remotecommand.TerminalSize) *termSizeQueue {
queue := &termSizeQueue{
incoming: make(chan remotecommand.TerminalSize, 1),
ctx: ctx,
}
queue.AddSize(initialSize)
return queue
}
func (r *termSizeQueue) Next() *remotecommand.TerminalSize {
select {
case <-r.ctx.Done():
return nil
case size := <-r.incoming:
return &size
}
}
func (r *termSizeQueue) AddSize(term remotecommand.TerminalSize) {
select {
case <-r.ctx.Done():
case r.incoming <- term:
}
}
func createKubeRestConfig(serverAddr, tlsServerName string, ca types.CertAuthority, clientCert, rsaKey []byte) (*rest.Config, error) {
var clusterCACerts [][]byte
for _, keyPair := range ca.GetTrustedTLSKeyPairs() {
clusterCACerts = append(clusterCACerts, keyPair.Cert)
}
return &rest.Config{
Host: serverAddr,
TLSClientConfig: rest.TLSClientConfig{
CertData: clientCert,
KeyData: rsaKey,
CAData: bytes.Join(clusterCACerts, []byte("\n")),
ServerName: tlsServerName,
},
}, nil
}
func (h *Handler) joinKubernetesSession(
ctx context.Context,
sessionID string,
mode types.SessionParticipantMode,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) error {
h.logger.InfoContext(ctx, "Attempting to join kubernetes existing session",
"session_id", sessionID,
"mode", mode,
"user", sctx.GetUser(),
)
if _, err := session.ParseID(sessionID); err != nil {
return trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return trace.Wrap(err)
}
tracker, err := clt.GetSessionTracker(ctx, sessionID)
if err != nil {
return trace.Wrap(err)
}
if tracker.GetSessionKind() != types.KubernetesSessionKind || tracker.GetState() == types.SessionState_SessionStateTerminated {
return trace.NotFound("Kubernetes session %v not found", sessionID)
}
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: session.Session{
Kind: types.KubernetesSessionKind,
ID: session.ID(tracker.GetName()),
Login: tracker.GetLogin(),
KubernetesClusterName: tracker.GetKubeCluster(),
ServerHostname: tracker.GetHostname(),
}})
if err != nil {
return trace.Wrap(err)
}
stream := terminal.NewStream(ctx, terminal.StreamConfig{
WS: ws,
Logger: h.logger,
// Disable all out of band handling of requests
Handlers: map[string]terminal.WSHandlerFunc{
defaults.WebsocketResize: func(ctx context.Context, envelope terminal.Envelope) {},
defaults.WebsocketFileTransferRequest: func(ctx context.Context, envelope terminal.Envelope) {},
defaults.WebsocketFileTransferDecision: func(ctx context.Context, envelope terminal.Envelope) {},
},
})
envelopeBytes, err := proto.Marshal(&terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketSessionMetadata,
Payload: string(sessionMetadataResponse),
})
if err != nil {
return trace.Wrap(err)
}
if err := stream.WriteMessage(websocket.BinaryMessage, envelopeBytes); err != nil {
return trace.Wrap(err)
}
authAccessPoint, err := cluster.CachingAccessPoint()
if err != nil {
return trace.Wrap(err)
}
netConfig, err := authAccessPoint.GetClusterNetworkingConfig(ctx)
if err != nil {
return trace.Wrap(err)
}
kubeAddr, tlsServerName, err := h.getKubeExecClusterData(netConfig)
if err != nil {
return trace.Wrap(err)
}
pk, err := keys.ParsePrivateKey(sctx.cfg.Session.GetTLSPriv())
if err != nil {
return trace.Wrap(err, "failed getting user private key from the session")
}
privateKeyPEM, err := pk.SoftwarePrivateKeyPEM()
if err != nil {
return trace.Wrap(err, "failed getting software private key")
}
publicKeyPEM, err := keys.MarshalPublicKey(pk.Public())
if err != nil {
return trace.Wrap(err, "failed to marshal public key")
}
certsReq := clientproto.UserCertsRequest{
TLSPublicKey: publicKeyPEM,
Username: sctx.GetUser(),
Expires: sctx.cfg.Session.GetExpiryTime(),
Format: constants.CertificateFormatStandard,
RouteToCluster: tracker.GetClusterName(),
KubernetesCluster: tracker.GetKubeCluster(),
Usage: clientproto.UserCertsRequest_Kubernetes,
}
var certs *clientproto.Certs
result, err := client.PerformSessionMFACeremony(ctx, client.PerformSessionMFACeremonyParams{
CurrentAuthClient: clt,
RootAuthClient: sctx.cfg.RootClient,
MFACeremony: newMFACeremony(stream.WSStream, sctx.cfg.RootClient.CreateAuthenticateChallenge, h.cfg.PublicProxyAddr),
MFAAgainstRoot: sctx.cfg.RootClusterName == tracker.GetClusterName(),
MFARequiredReq: &clientproto.IsMFARequiredRequest{
Target: &clientproto.IsMFARequiredRequest_KubernetesCluster{KubernetesCluster: tracker.GetKubeCluster()},
},
CertsReq: &certsReq,
})
if err != nil && !errors.Is(err, services.ErrSessionMFANotRequired) {
return trace.Wrap(err, "failed performing mfa ceremony")
} else if result != nil {
certs = result.NewCerts
}
if certs == nil {
certs, err = sctx.cfg.RootClient.GenerateUserCerts(ctx, certsReq)
if err != nil {
return trace.Wrap(err, "failed issuing user certs")
}
}
hostCA, err := h.auth.accessPoint.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: h.auth.clusterName,
}, false)
if err != nil {
return trace.Wrap(err)
}
certPool := x509.NewCertPool()
for _, keyPair := range hostCA.GetTrustedTLSKeyPairs() {
if ok := certPool.AppendCertsFromPEM(keyPair.Cert); !ok {
return trace.BadParameter("invalid ca format")
}
}
session, err := client.NewKubeSession(ctx,
client.KubeSessionConfig{
KubeProxyAddr: strings.TrimPrefix(kubeAddr, "https://"),
WebProxyAddr: h.cfg.ProxyWebAddr.String(),
TLSRoutingConnUpgradeRequired: netConfig.GetProxyListenerMode() == types.ProxyListenerMode_Multiplex,
EnableEscapeSequences: true,
Tracker: tracker,
TLSConfig: &tls.Config{
RootCAs: certPool,
ServerName: tlsServerName,
GetClientCertificate: func(*tls.CertificateRequestInfo) (*tls.Certificate, error) {
cert, err := tls.X509KeyPair(certs.TLS, privateKeyPEM)
if err != nil {
return nil, trace.Wrap(err)
}
return &cert, nil
},
},
Mode: mode,
AuthClient: func(ctx context.Context) (authclient.ClientI, error) {
return noopAuthClientCloser{clt}, nil
},
Ceremony: newMFACeremony(stream.WSStream, nil, h.cfg.PublicProxyAddr),
Stdin: stream,
Stdout: stream,
Stderr: stderrWriter{stream: stream},
})
if err != nil {
return trace.Wrap(err)
}
session.Wait()
if err := session.Detach(); err != nil && !terminal.IsOKWebsocketCloseError(err) {
return trace.Wrap(err)
}
return nil
}
type noopAuthClientCloser struct {
authclient.ClientI
}
func (noopAuthClientCloser) Close() error { return nil }
/*
* *
* * Teleport
* * Copyright (C) 2024 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"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"github.com/gravitational/teleport/lib/utils/diagnostics/latency"
)
// monitorLatency implements the Web UI's latency detector.
// It runs as long as the provided context has not expired.
//
// The latency of the provided websocket is monitored automatically,
// and the latency to the target endpoint is monitored with the provided pinger.
// The results of the latency calculation are reported to the web UI
// with the provided reporter.
func monitorLatency(
ctx context.Context,
clock clockwork.Clock,
ws latency.WebSocket,
endpointPinger latency.Pinger,
reporter latency.Reporter,
) error {
wsPinger, err := latency.NewWebsocketPinger(clock, ws)
if err != nil {
return trace.Wrap(err, "creating websocket pinger")
}
monitor, err := latency.NewMonitor(latency.MonitorConfig{
ClientPinger: wsPinger,
ServerPinger: endpointPinger,
Reporter: reporter,
Clock: clock,
})
if err != nil {
return trace.Wrap(err, "creating latency monitor")
}
monitor.Run(ctx)
return nil
}
// 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 (
"cmp"
"context"
"fmt"
"net/http"
"strconv"
"strings"
"time"
"github.com/coreos/go-semver/semver"
yaml "github.com/ghodss/yaml"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"google.golang.org/protobuf/types/known/durationpb"
"google.golang.org/protobuf/types/known/fieldmaskpb"
"github.com/gravitational/teleport/api/constants"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
machineidv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
tslices "github.com/gravitational/teleport/lib/utils/slices"
)
const (
// webUIFlowBotGitHubActionsSSH is the value of the webUIFlowLabelKey
// added to a resource created via the Bot GitHub Actions web UI flow.
webUIFlowBotGitHubActionsSSH = "github-actions-ssh"
)
type ListBotsResponse struct {
// Items is a list of resources retrieved.
Items []*machineidv1.Bot `json:"items"`
}
type CreateBotRequest struct {
// BotName is the name of the bot
BotName string `json:"botName"`
// Roles are the roles that the bot will be able to impersonate
Roles []string `json:"roles"`
// Traits are the traits that will be associated with the bot for the purposes of role
// templating.
// Where multiple specified with the same name, these will be merged by the
// server.
Traits []*machineidv1.Trait `json:"traits"`
}
// listBots returns a list of bots for a given cluster site. It does not leverage pagination from the UI. Due to the
// nature of the bot:user relationship, pagination is not yet supported. This endpoint will return all bots.
func (h *Handler) listBots(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
var items []*machineidv1.Bot
for pageToken := ""; ; {
bots, err := clt.BotServiceClient().ListBots(r.Context(), machineidv1.ListBotsRequest_builder{
PageSize: int32(1000),
PageToken: pageToken,
}.Build())
// todo (michellescripts) consider returning partial results
if err != nil {
return nil, trace.Wrap(err, "error getting bots")
}
items = append(items, bots.GetBots()...)
pageToken = bots.GetNextPageToken()
if pageToken == "" {
break
}
}
return ListBotsResponse{
Items: items,
}, nil
}
// createBot creates a bot
func (h *Handler) createBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var req *CreateBotRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
_, err = clt.BotServiceClient().CreateBot(r.Context(), machineidv1.CreateBotRequest_builder{
Bot: machineidv1.Bot_builder{
Kind: types.KindBot,
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: req.BotName,
Labels: map[string]string{
webUIFlowLabelKey: webUIFlowBotGitHubActionsSSH,
},
}.Build(),
Spec: machineidv1.BotSpec_builder{
Roles: req.Roles,
Traits: req.Traits,
}.Build(),
}.Build(),
}.Build())
if err != nil {
return nil, trace.Wrap(err, "error creating bot")
}
return OK(), nil
}
func (h *Handler) deleteBot(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
name := params.ByName("name")
if name == "" {
return nil, trace.BadParameter("missing bot name")
}
_, err = clt.BotServiceClient().DeleteBot(r.Context(), machineidv1.DeleteBotRequest_builder{BotName: name}.Build())
if err != nil {
return nil, trace.Wrap(err, "error deleting bot")
}
return OK(), nil
}
// CreateBotJoinTokenRequest represents a client request to
// create a bot join token
type CreateBotJoinTokenRequest struct {
// IntegrationName is the name attributed to the bot integration, which
// is used to name the resources created during the UI flow.
IntegrationName string `json:"integrationName"`
// JoinMethod is the joining method required in order to use this token.
JoinMethod types.JoinMethod `json:"joinMethod"`
// GitHub allows the configuration of options specific to the "github" join method.
GitHub *types.ProvisionTokenSpecV2GitHub `json:"gitHub"`
// WebFlowLabel is the value of the label attributed to bots created via the web UI
WebFlowLabel string `json:"webFlowLabel"`
}
// createBotJoinToken creates a bot join token
func (h *Handler) createBotJoinToken(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var req *CreateBotJoinTokenRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := types.ValidateJoinMethod(req.JoinMethod); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
spec := types.ProvisionTokenSpecV2{
Roles: []types.SystemRole{types.RoleBot},
JoinMethod: req.JoinMethod,
GitHub: req.GitHub,
BotName: req.IntegrationName,
}
provisionToken, err := types.NewProvisionTokenFromSpec(req.IntegrationName, time.Time{}, spec)
if err != nil {
return nil, trace.Wrap(err)
}
provisionToken.SetLabels(map[string]string{
webUIFlowLabelKey: req.WebFlowLabel,
})
err = clt.CreateToken(r.Context(), provisionToken)
if err != nil {
return nil, trace.Wrap(err, "error creating join token")
}
return &nodeJoinToken{
ID: provisionToken.GetName(),
Expiry: provisionToken.Expiry(),
Method: provisionToken.GetJoinMethod(),
}, nil
}
// getBot retrieves a bot by name
func (h *Handler) getBot(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
botName := p.ByName("name")
if botName == "" {
return nil, trace.BadParameter("empty name")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
bot, err := clt.BotServiceClient().GetBot(r.Context(), machineidv1.GetBotRequest_builder{
BotName: botName,
}.Build())
if err != nil {
return nil, trace.Wrap(err, "error querying bot")
}
return bot, nil
}
// updateBot updates a bot with provided roles. The only supported change via this endpoint today is roles.
// TODO(nicholasmarais1158) DELETE IN v20.0.0 - replaced by updateBotV2
// MUST delete with related code found in `web/packages/teleport/src/services/bot/bot.ts`
func (h *Handler) updateBotV1(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var request updateBotRequestV1
if err := httplib.ReadResourceJSON(r, &request); err != nil {
return nil, trace.Wrap(err)
}
return updateBot(r.Context(), p.ByName("name"), updateBotRequestV3{
Roles: request.Roles,
}, sctx, cluster)
}
type updateBotRequestV1 struct {
Roles []string `json:"roles"`
}
// updateBotV2 updates a bot with provided roles, traits and max_session_ttl.
// TODO(nicholasmarais1158) DELETE IN v20.0.0 - replaced by updateBotV3
// MUST delete with related code found in `web/packages/teleport/src/services/bot/bot.ts`
func (h *Handler) updateBotV2(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var request updateBotRequestV2
if err := httplib.ReadResourceJSON(r, &request); err != nil {
return nil, trace.Wrap(err)
}
return updateBot(r.Context(), p.ByName("name"), updateBotRequestV3{
Roles: request.Roles,
Traits: request.Traits,
MaxSessionTtl: request.MaxSessionTtl,
}, sctx, cluster)
}
type updateBotRequestV2 struct {
Roles []string `json:"roles"`
Traits []updateBotRequestTrait `json:"traits"`
MaxSessionTtl string `json:"max_session_ttl"`
}
type updateBotRequestTrait struct {
Name string `json:"name"`
Values []string `json:"values"`
}
// updateBot updates a bot with provided roles, traits, max_session_ttl and
// description.
func (h *Handler) updateBotV3(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var request updateBotRequestV3
if err := httplib.ReadResourceJSON(r, &request); err != nil {
return nil, trace.Wrap(err)
}
return updateBot(r.Context(), p.ByName("name"), request, sctx, cluster)
}
type updateBotRequestV3 struct {
Roles []string `json:"roles"`
Traits []updateBotRequestTrait `json:"traits"`
MaxSessionTtl string `json:"max_session_ttl"`
Description *string `json:"description"`
}
// updateBot updates a bot with provided roles, traits, max_session_ttl and
// description.
func updateBot(ctx context.Context, botName string, request updateBotRequestV3, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
if botName == "" {
return nil, trace.BadParameter("empty name")
}
mask, err := fieldmaskpb.New(&machineidv1.Bot{})
if err != nil {
return nil, trace.Wrap(err)
}
metadata := headerv1.Metadata_builder{
Name: botName,
}.Build()
spec := &machineidv1.BotSpec{}
if request.Roles != nil {
mask.Append(&machineidv1.Bot{}, "spec.roles")
spec.SetRoles(request.Roles)
}
if request.Traits != nil {
mask.Append(&machineidv1.Bot{}, "spec.traits")
traits := make([]*machineidv1.Trait, len(request.Traits))
for i, trait := range request.Traits {
traits[i] = machineidv1.Trait_builder{
Name: trait.Name,
Values: trait.Values,
}.Build()
}
spec.SetTraits(traits)
}
if request.MaxSessionTtl != "" {
mask.Append(&machineidv1.Bot{}, "spec.max_session_ttl")
ttl, err := time.ParseDuration(request.MaxSessionTtl)
if err != nil {
return nil, trace.Wrap(err)
}
spec.SetMaxSessionTtl(durationpb.New(ttl))
}
if request.Description != nil {
mask.Append(&machineidv1.Bot{}, "metadata.description")
metadata.SetDescription(*request.Description)
}
updated, err := clt.BotServiceClient().UpdateBot(ctx, machineidv1.UpdateBotRequest_builder{
UpdateMask: mask,
Bot: machineidv1.Bot_builder{
Kind: types.KindBot,
Version: types.V1,
Metadata: metadata,
Spec: spec,
}.Build(),
}.Build())
if err != nil {
return nil, trace.Wrap(err, "unable to find existing bot")
}
return updated, nil
}
// getBotInstance retrieves a bot instance by id
func (h *Handler) getBotInstance(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
botName := p.ByName("name")
instanceId := p.ByName("id")
if botName == "" {
return nil, trace.BadParameter("empty bot name")
}
if instanceId == "" {
return nil, trace.BadParameter("empty id")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
instance, err := clt.BotInstanceServiceClient().GetBotInstance(r.Context(), machineidv1.GetBotInstanceRequest_builder{
InstanceId: instanceId,
BotName: botName,
}.Build())
if err != nil {
return nil, trace.Wrap(err, "error querying bot instance")
}
yaml, err := yaml.Marshal(types.ProtoResource153ToLegacy(instance))
if err != nil {
return nil, trace.Wrap(err, "error stringifying to yaml")
}
return GetBotInstanceResponse{
BotInstance: instance,
YAML: string(yaml),
}, nil
}
type GetBotInstanceResponse struct {
BotInstance *machineidv1.BotInstance `json:"bot_instance"`
YAML string `json:"yaml"`
}
// listBotInstances returns a list of bot instances for a given cluster site.
func (h *Handler) listBotInstances(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
var pageSize int64 = 20
if r.URL.Query().Has("page_size") {
pageSize, err = strconv.ParseInt(r.URL.Query().Get("page_size"), 10, 32)
if err != nil {
return nil, trace.BadParameter("invalid page size")
}
}
var sort *types.SortBy
if r.URL.Query().Has("sort") {
sortString := r.URL.Query().Get("sort")
s := types.GetSortByFromString(sortString)
sort = &s
}
//nolint:staticcheck // SA1019. Kept for backward compatibility.
instances, err := clt.BotInstanceServiceClient().ListBotInstances(r.Context(), machineidv1.ListBotInstancesRequest_builder{
FilterBotName: r.URL.Query().Get("bot_name"),
PageSize: int32(pageSize),
PageToken: r.URL.Query().Get("page_token"),
FilterSearchTerm: r.URL.Query().Get("search"),
Sort: sort,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
uiInstances := tslices.Map(instances.GetBotInstances(), func(instance *machineidv1.BotInstance) BotInstance {
heartbeat := services.GetBotInstanceLatestHeartbeat(instance)
uiInstance := BotInstance{
InstanceId: instance.GetSpec().GetInstanceId(),
BotName: instance.GetSpec().GetBotName(),
}
if heartbeat != nil {
uiInstance.JoinMethodLatest = heartbeat.GetJoinMethod()
uiInstance.HostNameLatest = heartbeat.GetHostname()
uiInstance.VersionLatest = heartbeat.GetVersion()
uiInstance.ActiveAtLatest = heartbeat.GetRecordedAt().AsTime().Format(time.RFC3339)
uiInstance.OSLatest = heartbeat.GetOs()
}
return uiInstance
})
return ListBotInstancesResponse{
BotInstances: uiInstances,
NextPageToken: instances.GetNextPageToken(),
}, nil
}
// listBotInstancesV2 returns a list of bot instances for a given cluster site.
func (h *Handler) listBotInstancesV2(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
botName := r.URL.Query().Get("bot_name")
// Exhaustive view, per the scope_filter field docs.
var scopeFilter *scopesv1.Filter
if botName == "" {
scopeFilter = scopesv1.Filter_builder{Mode: scopesv1.Mode_MODE_ALL}.Build()
}
request := machineidv1.ListBotInstancesV2Request_builder{
PageToken: r.URL.Query().Get("page_token"),
SortField: r.URL.Query().Get("sort_field"),
Filter: machineidv1.ListBotInstancesV2Request_Filters_builder{
BotName: botName,
SearchTerm: r.URL.Query().Get("search"),
Query: r.URL.Query().Get("query"),
ScopeFilter: scopeFilter,
}.Build(),
}.Build()
if r.URL.Query().Has("page_size") {
pageSize, err := strconv.ParseInt(r.URL.Query().Get("page_size"), 10, 32)
if err != nil {
return nil, trace.BadParameter("invalid page size")
}
request.SetPageSize(int32(pageSize))
}
if r.URL.Query().Has("sort_dir") {
sortDir := r.URL.Query().Get("sort_dir")
request.SetSortDesc(strings.ToLower(sortDir) == "desc")
}
instances, err := clt.BotInstanceServiceClient().ListBotInstancesV2(r.Context(), request)
if err != nil {
return nil, trace.Wrap(err)
}
uiInstances := tslices.Map(instances.GetBotInstances(), func(instance *machineidv1.BotInstance) BotInstance {
heartbeat := services.GetBotInstanceLatestHeartbeat(instance)
authentication := services.GetBotInstanceLatestAuthentication(instance)
uiInstance := BotInstance{
InstanceId: instance.GetSpec().GetInstanceId(),
BotName: instance.GetSpec().GetBotName(),
}
if authentication != nil {
uiInstance.JoinMethodLatest = cmp.Or(
authentication.GetJoinAttrs().GetMeta().GetJoinMethod(),
authentication.GetJoinMethod(),
)
}
if heartbeat != nil {
uiInstance.HostNameLatest = heartbeat.GetHostname()
uiInstance.VersionLatest = heartbeat.GetVersion()
uiInstance.ActiveAtLatest = heartbeat.GetRecordedAt().AsTime().Format(time.RFC3339)
uiInstance.OSLatest = heartbeat.GetOs()
}
return uiInstance
})
return ListBotInstancesResponse{
BotInstances: uiInstances,
NextPageToken: instances.GetNextPageToken(),
}, nil
}
type ListBotInstancesResponse struct {
BotInstances []BotInstance `json:"bot_instances"`
NextPageToken string `json:"next_page_token,omitempty"`
}
type BotInstance struct {
InstanceId string `json:"instance_id"`
BotName string `json:"bot_name"`
JoinMethodLatest string `json:"join_method_latest,omitempty"`
HostNameLatest string `json:"host_name_latest,omitempty"`
VersionLatest string `json:"version_latest,omitempty"`
ActiveAtLatest string `json:"active_at_latest,omitempty"`
OSLatest string `json:"os_latest,omitempty"`
}
func (h *Handler) botInstanceMetrics(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
rsp := BotInstanceMetricsResponse{
RefreshAfterSeconds: int(constants.AutoUpdateBotInstanceReportPeriod.Seconds()),
}
// If no report is available yet, `UpgradeStatuses` will be nil.
report, err := clt.GetAutoUpdateBotInstanceReport(ctx)
switch {
case trace.IsNotFound(err):
return rsp, nil
case err != nil:
return nil, trace.Wrap(err)
}
// Our target version is the operator's selected auto-update tools version,
// or if there isn't one configured: the proxy version.
autoUpdateVersion, err := h.cfg.AccessPoint.GetAutoUpdateVersion(ctx)
if err != nil && !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
targetVersion, err := getToolsVersion(autoUpdateVersion)
if err != nil {
return nil, trace.Wrap(err)
}
// Returns the earliest possible version in a major release. It's based on:
// lib/utils.VersionBeforeAlpha.
lowerBound := func(major int64) semver.Version {
return semver.Version{Major: major, PreRelease: "aa"}
}
const versionField = "status.latest_heartbeat.version"
rsp.UpgradeStatuses = &BotInstanceUpgradeStatuses{
UpdatedAt: report.GetSpec().GetTimestamp().AsTime(),
UpToDate: BotInstanceUpgradeStatus{
Filter: fmt.Sprintf("%[1]s == %[2]q", versionField, targetVersion),
},
Unsupported: BotInstanceUpgradeStatus{
Filter: fmt.Sprintf(
"older_than(%[1]s, %[2]q) || %[1]s == %[3]q || newer_than(%[1]s, %[3]q)",
versionField,
lowerBound(targetVersion.Major-1),
lowerBound(targetVersion.Major+1),
),
},
PatchAvailable: BotInstanceUpgradeStatus{
Filter: fmt.Sprintf(
"between(%[1]s, %[2]q, %[3]q)",
versionField,
lowerBound(targetVersion.Major),
targetVersion,
),
},
RequiresUpgrade: BotInstanceUpgradeStatus{
Filter: fmt.Sprintf(
"between(%[1]s, %[2]q, %[3]q)",
versionField,
lowerBound(targetVersion.Major-1),
lowerBound(targetVersion.Major),
),
},
}
for _, groupMetrics := range report.GetSpec().GetGroups() {
for versionString, versionMetrics := range groupMetrics.GetVersions() {
version, err := semver.NewVersion(versionString)
if err != nil {
h.logger.ErrorContext(ctx,
"Failed to parse bot instance version string",
"version_string", versionString,
"error", err,
)
continue
}
switch {
case targetVersion.Equal(*version):
// Bot is up to date.
rsp.UpgradeStatuses.UpToDate.Count += int(versionMetrics.GetCount())
case targetVersion.LessThan(*version):
// Bot is running a newer version, we don't support this.
rsp.UpgradeStatuses.Unsupported.Count += int(versionMetrics.GetCount())
case targetVersion.Major == version.Major:
// Bot is running the right major version, but there's a minor
// or patch update available
rsp.UpgradeStatuses.PatchAvailable.Count += int(versionMetrics.GetCount())
case version.Major == targetVersion.Major-1:
// Bot is running the previous major version and should upgrade.
rsp.UpgradeStatuses.RequiresUpgrade.Count += int(versionMetrics.GetCount())
case version.Major < targetVersion.Major-1:
// Bot is running a version that is too old. In this case, the
// connection would be terminated so we shouldn't really see it.
rsp.UpgradeStatuses.Unsupported.Count += int(versionMetrics.GetCount())
default:
// The branches of this switch should be exhaustive, but just in case!
h.logger.DebugContext(ctx,
"Bot instance version comparison is missing a branch",
"bot_instance_version", version,
"target_version", targetVersion,
)
}
}
}
return rsp, nil
}
type BotInstanceMetricsResponse struct {
// RefreshAfterSeconds is the amount of time (in seconds) after receiving
// this response the client should poll for new metrics.
RefreshAfterSeconds int `json:"refresh_after_seconds"`
// UpgradeStatuses contains instance counts by "upgrade status".
UpgradeStatuses *BotInstanceUpgradeStatuses `json:"upgrade_statuses"`
}
type BotInstanceUpgradeStatuses struct {
// UpdatedAt is when these metrics were last updated.
UpdatedAt time.Time `json:"updated_at"`
// UpToDate means the instance matches the desired version.
UpToDate BotInstanceUpgradeStatus `json:"up_to_date"`
// Unsupported means the instance is running a release that is too old or
// too new for us to support.
Unsupported BotInstanceUpgradeStatus `json:"unsupported"`
// RequiresUpgrade means the instance is running a release from the previous
// major series. We can support it for now, but the next major upgrade will
// break compatibility.
RequiresUpgrade BotInstanceUpgradeStatus `json:"requires_upgrade"`
// PatchAvailable means the instance is running a release from the desired
// major series, but they're behind on a minor or patch release.
PatchAvailable BotInstanceUpgradeStatus `json:"patch_available"`
}
type BotInstanceUpgradeStatus struct {
// Count is the number of bot instances.
Count int `json:"count"`
// Filter is a predicate language filter that can be applied to the bot
// instance list to find matching instances.
Filter string `json:"filter"`
}
/*
* Teleport
* Copyright (C) 2026 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 (
"net/http"
"strings"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport"
autoupdatepb "github.com/gravitational/teleport/api/gen/proto/go/teleport/autoupdate/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/types/autoupdate"
"github.com/gravitational/teleport/api/utils/clientutils"
aur "github.com/gravitational/teleport/lib/autoupdate/report"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/web/ui"
)
// getManagedUpdatesDetails returns managed updates details.
func (h *Handler) getManagedUpdatesDetails(w http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
ctx := r.Context()
clt, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
response := &ui.ManagedUpdatesDetails{}
autoUpdateConfig, err := clt.GetAutoUpdateConfig(ctx)
if err != nil {
if !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
autoUpdateConfig = nil
}
autoUpdateVersion, err := clt.GetAutoUpdateVersion(ctx)
if err != nil {
if !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
autoUpdateVersion = nil
}
response.Tools = getToolsInfo(autoUpdateConfig, autoUpdateVersion)
rollout, err := clt.GetAutoUpdateAgentRollout(ctx)
if err != nil {
if !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
rollout = nil
}
if rollout != nil {
response.Rollout = getRolloutInfo(rollout)
}
reports, err := stream.Collect(clientutils.Resources(ctx, clt.ListAutoUpdateAgentReports))
if err != nil {
if !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
reports = nil
}
// Filter and aggregate version counts from the reports
validReports := aur.ValidReports(reports, time.Now())
versionCountsByGroup := aur.AggregateVersionCounts(validReports)
if rollout != nil {
response.Groups = getGroupsInfo(rollout, versionCountsByGroup)
response.OrphanedAgentVersionCounts = getOrphanedAgentCounts(rollout, versionCountsByGroup)
}
// Get cluster maintenance info if this is a cloud cluster
if features := h.GetClusterFeatures(); features.GetCloud() {
maintenanceConfig, err := clt.GetClusterMaintenanceConfig(ctx)
if err != nil {
if !trace.IsNotFound(err) {
return nil, trace.Wrap(err)
}
maintenanceConfig = nil
}
if maintenanceConfig != nil {
response.ClusterMaintenance = getClusterMaintenanceInfo(maintenanceConfig)
}
}
return response, nil
}
// getToolsInfo builds the ToolsAutoUpdateInfo object.
func getToolsInfo(config *autoupdatepb.AutoUpdateConfig, version *autoupdatepb.AutoUpdateVersion) *ui.ToolsAutoUpdateInfo {
var mode, targetVersion string
if config != nil {
mode = config.GetSpec().GetTools().GetMode()
}
if version != nil {
targetVersion = version.GetSpec().GetTools().GetTargetVersion()
}
// If empty, return nil
if mode == "" && targetVersion == "" {
return nil
}
return &ui.ToolsAutoUpdateInfo{
Mode: mode,
TargetVersion: targetVersion,
}
}
// getRolloutInfo builds the RolloutInfo object.
func getRolloutInfo(rollout *autoupdatepb.AutoUpdateAgentRollout) *ui.RolloutInfo {
if rollout == nil || rollout.GetSpec() == nil {
return nil
}
spec := rollout.GetSpec()
status := rollout.GetStatus()
info := &ui.RolloutInfo{
StartVersion: spec.GetStartVersion(),
TargetVersion: spec.GetTargetVersion(),
Strategy: spec.GetStrategy(),
Schedule: spec.GetSchedule(),
State: strings.ToLower(aur.UserFriendlyState(status.GetState())),
Mode: spec.GetAutoupdateMode(),
}
// Set the rollout start time
if status != nil {
if startTime := status.GetStartTime(); startTime != nil && startTime.IsValid() {
t := startTime.AsTime()
if !t.IsZero() && t.Unix() != 0 {
info.StartTime = &t
}
}
}
return info
}
// getGroupsInfo builds the list of RolloutGroupInfo objects.
func getGroupsInfo(rollout *autoupdatepb.AutoUpdateAgentRollout, versionCountsByGroup map[string]map[string]int) []ui.RolloutGroupInfo {
if rollout == nil {
return nil
}
groups := rollout.GetStatus().GetGroups()
if len(groups) == 0 {
return nil
}
out := make([]ui.RolloutGroupInfo, 0, len(groups))
for i, group := range groups {
groupInfo := ui.RolloutGroupInfo{
Name: group.GetName(),
State: strings.ToLower(aur.UserFriendlyState(group.GetState())),
InitialCount: group.GetInitialCount(),
PresentCount: group.GetPresentCount(),
UpToDateCount: group.GetUpToDateCount(),
StateReason: group.GetLastUpdateReason(),
CanaryCount: group.GetCanaryCount(),
IsCatchAll: i == len(groups)-1,
}
// Only set the position if the strategy is halt-on-error
if rollout.GetSpec().GetStrategy() == autoupdate.AgentsStrategyHaltOnError {
groupInfo.Position = i + 1
}
// Set the group start time
if startTime := group.GetStartTime(); startTime != nil && startTime.IsValid() {
t := startTime.AsTime()
if !t.IsZero() && t.Unix() != 0 {
groupInfo.StartTime = &t
}
}
// Set the last update time
if lastUpdateTime := group.GetLastUpdateTime(); lastUpdateTime != nil && lastUpdateTime.IsValid() {
t := lastUpdateTime.AsTime()
if !t.IsZero() && t.Unix() != 0 {
groupInfo.LastUpdateTime = &t
}
}
// Add the version counts from aggregated reports
if counts, ok := versionCountsByGroup[group.GetName()]; ok && len(counts) > 0 {
groupInfo.AgentVersionCounts = counts
}
// Calculate the CanarySuccessCount
if groupInfo.CanaryCount > 0 {
var successCount uint64
for _, canary := range group.GetCanaries() {
if canary.GetSuccess() {
successCount++
}
}
groupInfo.CanarySuccessCount = successCount
}
out = append(out, groupInfo)
}
return out
}
// getClusterMaintenaceInfo builds the ClusterMaintenanceInfo object.
func getClusterMaintenanceInfo(cmc types.ClusterMaintenanceConfig) *ui.ClusterMaintenanceInfo {
window, ok := cmc.GetAgentUpgradeWindow()
if !ok {
return nil
}
return &ui.ClusterMaintenanceInfo{
ControlPlaneVersion: teleport.Version,
MaintenanceWeekdays: window.Weekdays,
MaintenanceStartHour: int(window.UTCStartHour),
}
}
// getOrphanedAgentCounts returns version counts for agents reporting group names
// that don't match any defined rollout group.
func getOrphanedAgentCounts(rollout *autoupdatepb.AutoUpdateAgentRollout, versionCountsByGroup map[string]map[string]int) map[string]int {
if rollout == nil || len(versionCountsByGroup) == 0 {
return nil
}
// Get the defined rollout groups
definedGroups := make(map[string]bool)
for _, group := range rollout.GetStatus().GetGroups() {
definedGroups[group.GetName()] = true
}
// Calculate how many agents don't belong to any of those groups, and their version.
orphanedCounts := make(map[string]int)
for groupName, versionCounts := range versionCountsByGroup {
if !definedGroups[groupName] {
for version, count := range versionCounts {
orphanedCounts[version] += count
}
}
}
if len(orphanedCounts) == 0 {
return nil
}
return orphanedCounts
}
func getAutoUpdateServiceClient(sctx *SessionContext) autoupdatepb.AutoUpdateServiceClient {
return autoupdatepb.NewAutoUpdateServiceClient(sctx.GetClientConnection())
}
// startGroupUpdate starts an update for a specified rollout group.
func (h *Handler) startGroupUpdate(w http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
ctx := r.Context()
groupName := params.ByName("groupName")
if groupName == "" {
return nil, trace.BadParameter("group name is required")
}
var req ui.StartGroupUpdateRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
state := autoupdatepb.AutoUpdateAgentGroupState_AUTO_UPDATE_AGENT_GROUP_STATE_UNSPECIFIED
// If the force flag is set to true, set the desired state to active to skip canary phase.
if req.Force {
state = autoupdatepb.AutoUpdateAgentGroupState_AUTO_UPDATE_AGENT_GROUP_STATE_ACTIVE
}
client := getAutoUpdateServiceClient(sctx)
rollout, err := client.TriggerAutoUpdateAgentGroup(ctx, autoupdatepb.TriggerAutoUpdateAgentGroupRequest_builder{
Groups: []string{groupName},
DesiredState: state,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
group, err := findGroupInfo(rollout, groupName)
if err != nil {
return nil, trace.Wrap(err)
}
return &ui.GroupActionResponse{Group: group}, nil
}
// markGroupDone marks a specified rollout group as done.
func (h *Handler) markGroupDone(w http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
ctx := r.Context()
groupName := params.ByName("groupName")
if groupName == "" {
return nil, trace.BadParameter("group name is required")
}
client := getAutoUpdateServiceClient(sctx)
rollout, err := client.ForceAutoUpdateAgentGroup(ctx, autoupdatepb.ForceAutoUpdateAgentGroupRequest_builder{
Groups: []string{groupName},
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
group, err := findGroupInfo(rollout, groupName)
if err != nil {
return nil, trace.Wrap(err)
}
return &ui.GroupActionResponse{Group: group}, nil
}
// rollbackGroup rolls back a specified rollout group.
func (h *Handler) rollbackGroup(w http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
ctx := r.Context()
groupName := params.ByName("groupName")
if groupName == "" {
return nil, trace.BadParameter("group name is required")
}
auClient := getAutoUpdateServiceClient(sctx)
rollout, err := auClient.RollbackAutoUpdateAgentGroup(ctx, autoupdatepb.RollbackAutoUpdateAgentGroupRequest_builder{
Groups: []string{groupName},
AllStartedGroups: false,
}.Build())
if err != nil {
return nil, trace.Wrap(err)
}
group, err := findGroupInfo(rollout, groupName)
if err != nil {
return nil, trace.Wrap(err)
}
return &ui.GroupActionResponse{Group: group}, nil
}
// findGroupInfo gets the RolloutGroupInfo for a specified group name.
func findGroupInfo(rollout *autoupdatepb.AutoUpdateAgentRollout, groupName string) (*ui.RolloutGroupInfo, error) {
if rollout == nil || rollout.GetStatus() == nil {
return nil, trace.NotFound("group %q not found in rollout", groupName)
}
groups := getGroupsInfo(rollout, nil)
for i := range groups {
if groups[i].Name == groupName {
return &groups[i], nil
}
}
return nil, trace.NotFound("group %q not found in rollout", groupName)
}
/*
* 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"
"net/http"
"net/url"
"strings"
"github.com/google/uuid"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
mfav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v1"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/client/sso"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
// getMFADevicesWithTokenHandle retrieves the list of registered MFA devices for the user defined in token.
func (h *Handler) getMFADevicesWithTokenHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
mfas, err := h.cfg.ProxyClient.GetMFADevices(r.Context(), &proto.GetMFADevicesRequest{
TokenID: p.ByName("token"),
})
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeMFADevices(mfas.GetDevices()), nil
}
// getMFADevicesHandle retrieves the list of registered MFA devices for the user in context (logged in user).
func (h *Handler) getMFADevicesHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params, c *SessionContext) (any, error) {
clt, err := c.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
mfas, err := clt.GetMFADevices(r.Context(), &proto.GetMFADevicesRequest{})
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeMFADevices(mfas.GetDevices()), nil
}
// deleteMFADeviceWithTokenHandle deletes a mfa device for the user defined in the `token`, given as a query parameter.
func (h *Handler) deleteMFADeviceWithTokenHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
if err := h.GetProxyClient().DeleteMFADeviceSync(r.Context(), &proto.DeleteMFADeviceSyncRequest{
TokenID: p.ByName("token"),
DeviceName: p.ByName("devicename"),
}); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
type addMFADeviceRequest struct {
// PrivilegeTokenID is privilege token id.
PrivilegeTokenID string `json:"tokenId"`
// DeviceName is the name of new mfa device.
DeviceName string `json:"deviceName"`
// SecondFactorToken is the totp code.
SecondFactorToken string `json:"secondFactorToken"`
// WebauthnRegisterResponse is a WebAuthn registration challenge response.
WebauthnRegisterResponse *wantypes.CredentialCreationResponse `json:"webauthnRegisterResponse"`
// DeviceUsage is the intended usage of the device (MFA, Passwordless, etc).
// It mimics the proto.DeviceUsage enum.
// Defaults to MFA.
DeviceUsage string `json:"deviceUsage"`
}
// addMFADeviceHandle adds a new mfa device for the user defined in the token.
func (h *Handler) addMFADeviceHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
var req addMFADeviceRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
deviceUsage, err := getDeviceUsage(req.DeviceUsage)
if err != nil {
return nil, trace.Wrap(err)
}
protoReq := &proto.AddMFADeviceSyncRequest{
TokenID: req.PrivilegeTokenID,
NewDeviceName: req.DeviceName,
DeviceUsage: deviceUsage,
}
switch {
case req.SecondFactorToken != "":
protoReq.NewMFAResponse = &proto.MFARegisterResponse{Response: &proto.MFARegisterResponse_TOTP{
TOTP: &proto.TOTPRegisterResponse{Code: req.SecondFactorToken},
}}
case req.WebauthnRegisterResponse != nil:
protoReq.NewMFAResponse = &proto.MFARegisterResponse{Response: &proto.MFARegisterResponse_Webauthn{
Webauthn: wantypes.CredentialCreationResponseToProto(req.WebauthnRegisterResponse),
}}
default:
return nil, trace.BadParameter("missing new mfa credentials")
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
if _, err := clt.AddMFADeviceSync(r.Context(), protoReq); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
type CreateAuthenticateChallengeRequest struct {
IsMFARequiredRequest *IsMFARequiredRequest `json:"is_mfa_required_req"`
ChallengeScope int `json:"challenge_scope"`
ChallengeAllowReuse bool `json:"challenge_allow_reuse"`
UserVerificationRequirement string `json:"user_verification_requirement"`
ProxyAddress string `json:"proxy_address"`
BrowserMFARequestID string `json:"browser_mfa_request_id"`
}
// createAuthenticateChallengeHandle creates and returns MFA authentication challenges for the user in context (logged in user).
// Used when users need to re-authenticate their second factors.
func (h *Handler) createAuthenticateChallengeHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params, c *SessionContext) (any, error) {
ctx := r.Context()
var req CreateAuthenticateChallengeRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := c.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
var mfaRequiredCheckProto *proto.IsMFARequiredRequest
if req.IsMFARequiredRequest != nil {
mfaRequiredCheckProto, err = h.checkAndGetProtoRequest(ctx, c, req.IsMFARequiredRequest)
if err != nil {
return nil, trace.Wrap(err)
}
// If this is an mfa required check for a leaf host, we need to check the requirement through
// the leaf cluster, rather than through root in the authenticate challenge request below
//
// TODO(Joerger): Currently, the only leafs hosts that we check mfa requirements for directly
// are apps. If we need to check other hosts directly, rather than through websocket flow,
// we'll need to include their clusterID in the request like we do for apps.
appReq := mfaRequiredCheckProto.GetApp()
if appReq != nil && appReq.ClusterName != c.cfg.RootClusterName {
site, err := h.getSiteByClusterName(ctx, c, appReq.ClusterName)
if err != nil {
return nil, trace.Wrap(err)
}
clusterClient, err := c.GetUserClient(ctx, site)
if err != nil {
return false, trace.Wrap(err)
}
res, err := clusterClient.IsMFARequired(ctx, mfaRequiredCheckProto)
if err != nil {
return false, trace.Wrap(err)
}
if !res.Required {
return &client.MFAAuthenticateChallenge{}, nil
}
// We don't want to check again through the root cluster below.
mfaRequiredCheckProto = nil
}
}
allowReuse := mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_NO
if req.ChallengeAllowReuse {
allowReuse = mfav1.ChallengeAllowReuse_CHALLENGE_ALLOW_REUSE_YES
}
// Prepare an sso client redirect URL in case the user has an SSO MFA device.
ssoClientRedirectURL, err := url.Parse(sso.WebMFARedirect)
if err != nil {
return nil, trace.Wrap(err)
}
// id is used by the front end to differentiate between separate ongoing SSO challenges.
id, err := uuid.NewRandom()
if err != nil {
return nil, trace.Wrap(err)
}
channelID := id.String()
query := ssoClientRedirectURL.Query()
query.Set("channel_id", channelID)
ssoClientRedirectURL.RawQuery = query.Encode()
// If BrowserMFARequestID is set, don't set the challenge extensions.
// They will be gotten from the stored MFASession on the backend and
// applied to the challenge.
var challengeExtensions *mfav1.ChallengeExtensions
if req.BrowserMFARequestID == "" {
challengeExtensions = &mfav1.ChallengeExtensions{
Scope: mfav1.ChallengeScope(req.ChallengeScope),
AllowReuse: allowReuse,
UserVerificationRequirement: req.UserVerificationRequirement,
}
}
chal, err := clt.CreateAuthenticateChallenge(ctx, &proto.CreateAuthenticateChallengeRequest{
Request: &proto.CreateAuthenticateChallengeRequest_ContextUser{
ContextUser: &proto.ContextUser{},
},
MFARequiredCheck: mfaRequiredCheckProto,
ChallengeExtensions: challengeExtensions,
SSOClientRedirectURL: ssoClientRedirectURL.String(),
ProxyAddress: req.ProxyAddress,
BrowserMFARequestID: req.BrowserMFARequestID,
})
if err != nil {
return nil, trace.Wrap(err)
}
return makeAuthenticateChallenge(chal, channelID), nil
}
// createAuthenticateChallengeWithTokenHandle creates and returns MFA authenticate challenges for the user defined in token.
func (h *Handler) createAuthenticateChallengeWithTokenHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
chal, err := h.cfg.ProxyClient.CreateAuthenticateChallenge(r.Context(), &proto.CreateAuthenticateChallengeRequest{
Request: &proto.CreateAuthenticateChallengeRequest_RecoveryStartTokenID{
RecoveryStartTokenID: p.ByName("token"),
},
ChallengeExtensions: &mfav1.ChallengeExtensions{
Scope: mfav1.ChallengeScope_CHALLENGE_SCOPE_ACCOUNT_RECOVERY,
},
SSOClientRedirectURL: sso.WebMFARedirect,
})
if err != nil {
return nil, trace.Wrap(err)
}
return makeAuthenticateChallenge(chal, "" /*channelID*/), nil
}
type createRegisterChallengeWithTokenRequest struct {
// DeviceType is the type of MFA device to get a register challenge for.
DeviceType string `json:"deviceType"`
// DeviceUsage is the intended usage of the device (MFA, Passwordless, etc).
// It mimics the proto.DeviceUsage enum.
// Defaults to MFA.
DeviceUsage string `json:"deviceUsage"`
}
// createRegisterChallengeWithTokenHandle creates and returns MFA register challenges for a new device for the specified device type.
func (h *Handler) createRegisterChallengeWithTokenHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req createRegisterChallengeWithTokenRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
var deviceType proto.DeviceType
switch req.DeviceType {
case "totp":
deviceType = proto.DeviceType_DEVICE_TYPE_TOTP
case "webauthn":
deviceType = proto.DeviceType_DEVICE_TYPE_WEBAUTHN
default:
return nil, trace.BadParameter("MFA device type %q unsupported", req.DeviceType)
}
deviceUsage, err := getDeviceUsage(req.DeviceUsage)
if err != nil {
return nil, trace.Wrap(err)
}
chal, err := h.cfg.ProxyClient.CreateRegisterChallenge(r.Context(), &proto.CreateRegisterChallengeRequest{
TokenID: p.ByName("token"),
DeviceType: deviceType,
DeviceUsage: deviceUsage,
})
if err != nil {
return nil, trace.Wrap(err)
}
resp := &client.MFARegisterChallenge{}
switch chal.GetRequest().(type) {
case *proto.MFARegisterChallenge_TOTP:
resp.TOTP = &client.TOTPRegisterChallenge{
QRCode: chal.GetTOTP().GetQRCode(),
}
case *proto.MFARegisterChallenge_Webauthn:
resp.Webauthn = wantypes.CredentialCreationFromProto(chal.GetWebauthn())
}
return resp, nil
}
func getDeviceUsage(reqUsage string) (proto.DeviceUsage, error) {
var deviceUsage proto.DeviceUsage
switch strings.ToLower(reqUsage) {
case "", "mfa":
deviceUsage = proto.DeviceUsage_DEVICE_USAGE_MFA
case "passwordless":
deviceUsage = proto.DeviceUsage_DEVICE_USAGE_PASSWORDLESS
default:
return proto.DeviceUsage_DEVICE_USAGE_UNSPECIFIED, trace.BadParameter("device usage %q unsupported", reqUsage)
}
return deviceUsage, nil
}
type isMFARequiredDatabase struct {
// ServiceName is the database service name.
ServiceName string `json:"service_name"`
// Protocol is the type of the database protocol
// eg: "postgres", "mysql", "mongodb", etc.
Protocol string `json:"protocol"`
// Username is an optional database username.
Username string `json:"username,omitempty"`
// DatabaseName is an optional database name.
DatabaseName string `json:"database_name,omitempty"`
}
type isMFARequiredKube struct {
// ClusterName is the name of the kube cluster.
ClusterName string `json:"cluster_name"`
}
type isMFARequiredNode struct {
// NodeName can be node's hostname or UUID.
NodeName string `json:"node_name"`
// Login is the OS login name.
Login string `json:"login"`
}
type isMFARequiredWindowsDesktop struct {
// DesktopName is the Windows Desktop server name.
DesktopName string `json:"desktop_name"`
// Login is the Windows desktop user login.
Login string `json:"login"`
}
type isMFARequiredLinuxDesktop struct {
// DesktopName is the Linux Desktop server name.
DesktopName string `json:"desktop_name"`
// Login is the Linux desktop user login.
Login string `json:"login"`
}
type IsMFARequiredApp struct {
// ResolveAppParams contains info used to resolve an application
ResolveAppParams
}
type isMFARequiredAdminAction struct{}
type IsMFARequiredRequest struct {
// Database contains fields required to check if target database
// requires MFA check.
Database *isMFARequiredDatabase `json:"database,omitempty"`
// Node contains fields required to check if target node
// requires MFA check.
Node *isMFARequiredNode `json:"node,omitempty"`
// WindowsDesktop contains fields required to check if target
// windows desktop requires MFA check.
WindowsDesktop *isMFARequiredWindowsDesktop `json:"windows_desktop,omitempty"`
// LinuxDesktop contains fields required to check if target
// linux desktop requires MFA check.
LinuxDesktop *isMFARequiredLinuxDesktop `json:"linux_desktop,omitempty"`
// Kube is the name of the kube cluster to check if target cluster
// requires MFA check.
Kube *isMFARequiredKube `json:"kube,omitempty"`
// App contains fields required to resolve an application and check if
// the target application requires MFA check.
App *IsMFARequiredApp `json:"app,omitempty"`
// AdminAction is the name of the admin action RPC to check if MFA is required.
AdminAction *isMFARequiredAdminAction `json:"admin_action,omitempty"`
}
func (h *Handler) checkAndGetProtoRequest(ctx context.Context, scx *SessionContext, r *IsMFARequiredRequest) (*proto.IsMFARequiredRequest, error) {
numRequests := 0
var protoReq *proto.IsMFARequiredRequest
if r.Database != nil {
numRequests++
if r.Database.ServiceName == "" {
return nil, trace.BadParameter("missing service_name for checking database target")
}
if r.Database.Protocol == "" {
return nil, trace.BadParameter("missing protocol for checking database target")
}
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_Database{
Database: &proto.RouteToDatabase{
ServiceName: r.Database.ServiceName,
Protocol: r.Database.Protocol,
Database: r.Database.DatabaseName,
Username: r.Database.Username,
},
},
}
}
if r.Kube != nil {
numRequests++
if r.Kube.ClusterName == "" {
return nil, trace.BadParameter("missing cluster_name for checking kubernetes cluster target")
}
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_KubernetesCluster{
KubernetesCluster: r.Kube.ClusterName,
},
}
}
if r.WindowsDesktop != nil {
numRequests++
if r.WindowsDesktop.DesktopName == "" {
return nil, trace.BadParameter("missing desktop_name for checking windows desktop target")
}
if r.WindowsDesktop.Login == "" {
return nil, trace.BadParameter("missing login for checking windows desktop target")
}
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_WindowsDesktop{
WindowsDesktop: &proto.RouteToWindowsDesktop{
WindowsDesktop: r.WindowsDesktop.DesktopName,
Login: r.WindowsDesktop.Login,
},
},
}
}
if r.LinuxDesktop != nil {
numRequests++
if r.LinuxDesktop.DesktopName == "" {
return nil, trace.BadParameter("missing desktop_name for checking linux desktop target")
}
if r.LinuxDesktop.Login == "" {
return nil, trace.BadParameter("missing login for checking linux desktop target")
}
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_LinuxDesktop{
LinuxDesktop: &proto.RouteToLinuxDesktop{
LinuxDesktop: r.LinuxDesktop.DesktopName,
Login: r.LinuxDesktop.Login,
},
},
}
}
if r.Node != nil {
numRequests++
if r.Node.Login == "" {
return nil, trace.BadParameter("missing login for checking node target")
}
if r.Node.NodeName == "" {
return nil, trace.BadParameter("missing node_name for checking node target")
}
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_Node{
Node: &proto.NodeLogin{
Login: r.Node.Login,
Node: r.Node.NodeName,
},
},
}
}
if r.App != nil {
resolvedApp, err := h.resolveApp(ctx, scx, r.App.ResolveAppParams)
if err != nil {
return nil, trace.Wrap(err, "unable to resolve FQDN: %v", r.App.FQDNHint)
}
numRequests++
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_App{
App: &proto.RouteToApp{
Name: resolvedApp.App.GetName(),
PublicAddr: resolvedApp.App.GetPublicAddr(),
ClusterName: resolvedApp.ClusterName,
},
},
}
}
if r.AdminAction != nil {
numRequests++
protoReq = &proto.IsMFARequiredRequest{
Target: &proto.IsMFARequiredRequest_AdminAction{
AdminAction: &proto.AdminAction{},
},
}
}
if numRequests > 1 {
return nil, trace.BadParameter("only one target is allowed for MFA check")
}
if protoReq == nil {
return nil, trace.BadParameter("missing target for MFA check")
}
return protoReq, nil
}
type isMfaRequiredResponse struct {
Required bool `json:"required"`
}
// isMFARequired is the [ClusterHandler] implementer for checking if MFA is required for a given target.
func (h *Handler) isMFARequired(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
var httpReq *IsMFARequiredRequest
if err := httplib.ReadResourceJSON(r, &httpReq); err != nil {
return nil, trace.Wrap(err)
}
required, err := h.checkMFARequired(r.Context(), httpReq, sctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
return isMfaRequiredResponse{Required: required}, nil
}
// checkMFARequired checks if MFA is required for the target specified in the [isMFARequiredRequest].
func (h *Handler) checkMFARequired(ctx context.Context, req *IsMFARequiredRequest, sctx *SessionContext, cluster reversetunnelclient.Cluster) (bool, error) {
protoReq, err := h.checkAndGetProtoRequest(ctx, sctx, req)
if err != nil {
return false, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return false, trace.Wrap(err)
}
res, err := clt.IsMFARequired(ctx, protoReq)
if err != nil {
return false, trace.Wrap(err)
}
return res.GetRequired(), nil
}
// makeAuthenticateChallenge converts proto to JSON format.
func makeAuthenticateChallenge(protoChal *proto.MFAAuthenticateChallenge, ssoChannelID string) *client.MFAAuthenticateChallenge {
chal := &client.MFAAuthenticateChallenge{
TOTPChallenge: protoChal.GetTOTP() != nil,
}
if protoChal.GetWebauthnChallenge() != nil {
chal.WebauthnChallenge = wantypes.CredentialAssertionFromProto(protoChal.WebauthnChallenge)
}
if protoChal.GetSSOChallenge() != nil {
chal.SSOChallenge = client.SSOChallengeFromProto(protoChal.GetSSOChallenge())
chal.SSOChallenge.ChannelID = ssoChannelID
}
if protoChal.GetBrowserMFAChallenge() != nil {
chal.BrowserMFAChallenge = client.BrowserChallengeFromProto(protoChal.GetBrowserMFAChallenge())
}
return chal
}
/*
* 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 (
"encoding/json"
proto "github.com/gogo/protobuf/proto"
"github.com/gravitational/trace"
authproto "github.com/gravitational/teleport/api/client/proto"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/web/mfajson"
"github.com/gravitational/teleport/lib/web/terminal"
)
// protobufMFACodec converts MFA challenges and responses to the protobuf
// format used by SSH web sessions
type protobufMFACodec struct{}
func (protobufMFACodec) Encode(chal *client.MFAAuthenticateChallenge, envelopeType string) ([]byte, error) {
jsonBytes, err := json.Marshal(chal)
if err != nil {
return nil, trace.Wrap(err)
}
envelope := &terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: envelopeType,
Payload: string(jsonBytes),
}
protoBytes, err := proto.Marshal(envelope)
if err != nil {
return nil, trace.Wrap(err)
}
return protoBytes, nil
}
func (protobufMFACodec) DecodeResponse(bytes []byte, envelopeType string) (*authproto.MFAAuthenticateResponse, error) {
return mfajson.Decode(bytes, envelopeType)
}
func (protobufMFACodec) DecodeChallenge(bytes []byte, envelopeType string) (*authproto.MFAAuthenticateChallenge, error) {
var challenge client.MFAAuthenticateChallenge
if err := json.Unmarshal(bytes, &challenge); err != nil {
return nil, trace.Wrap(err)
}
return &authproto.MFAAuthenticateChallenge{
WebauthnChallenge: wantypes.CredentialAssertionToProto(challenge.WebauthnChallenge),
}, nil
}
/*
* Teleport
* Copyright (C) 2025 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 (
"bytes"
"embed"
"encoding/json"
"fmt"
"net/http"
template "github.com/DataDog/datadog-agent/pkg/template/text"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport"
headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1"
machineidv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/tfgen"
"github.com/gravitational/teleport/lib/tfgen/transform"
"github.com/gravitational/teleport/lib/utils/slices"
)
// machineIDWizardGenerateIaC generates IaC code for the Machine Identity CI/CD wizards.
func (h *Handler) machineIDWizardGenerateIaC(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, _ *SessionContext, _ reversetunnelclient.Cluster) (any, error) {
var req machineIDWizardGenerateIaCRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
return nil, trace.Wrap(err)
}
switch {
// We currently only support deploying from GitHub Actions to Kubernetes but
// this endpoint will support other sources and destinations in the future.
case req.SourceType != "github":
return nil, trace.BadParameter("source_type must be one of: [github]")
case req.DestinationType != "kubernetes":
return nil, trace.BadParameter("destination_type must be one of: [kubernetes]")
// We also currently only support a single repository, but we may support
// multiple in the future, so the allow field is a slice.
case req.GitHub != nil && len(req.GitHub.Allow) != 1:
return nil, trace.BadParameter("github.allow: must contain exactly one item")
}
if req.GitHub == nil {
// Default to *something* so the generated code isn't completely broken.
req.GitHub = &machineIDWizardRequestGitHub{
Allow: []machineIDWizardRequestGitHubAllow{
{
Repository: "repository",
Owner: "organization",
},
},
}
}
namePrefix := fmt.Sprintf(
"gha-%s-%s",
req.GitHub.Allow[0].Owner,
req.GitHub.Allow[0].Repository,
)
// Role resource.
role := &types.RoleV6{
Kind: types.KindRole,
Version: types.V7,
Metadata: types.Metadata{
Name: fmt.Sprintf("%s-kube-access", namePrefix),
},
Spec: types.RoleSpecV6{
Allow: types.RoleConditions{},
},
}
var roleOpts []tfgen.GenerateOpt
if req.Kubernetes != nil {
role.Spec.Allow.KubernetesLabels = req.Kubernetes.Labels
role.Spec.Allow.KubernetesResources = req.Kubernetes.Resources
role.Spec.Allow.KubeGroups = req.Kubernetes.Groups
role.Spec.Allow.KubeUsers = req.Kubernetes.Users
} else {
roleOpts = append(roleOpts, tfgen.WithFieldComment("spec.allow.kubernetes_labels", "kubernetes_labels will be added in the next step."))
roleOpts = append(roleOpts, tfgen.WithFieldComment("spec.allow.kubernetes_groups", "kubernetes_groups will be added in the next step."))
}
// Bot resource.
bot := machineidv1.Bot_builder{
Kind: types.KindBot,
Version: types.V1,
Metadata: headerv1.Metadata_builder{
Name: namePrefix,
}.Build(),
Spec: machineidv1.BotSpec_builder{
Roles: []string{role.GetName()},
}.Build(),
}.Build()
botOpts := []tfgen.GenerateOpt{
tfgen.WithFieldTransform("spec.traits", transform.BotTraits),
}
// Join token resource.
token := &types.ProvisionTokenV2{
Kind: types.KindToken,
Version: types.V2,
Metadata: types.Metadata{
Name: namePrefix,
},
Spec: types.ProvisionTokenSpecV2{
Roles: []types.SystemRole{types.RoleBot},
JoinMethod: types.JoinMethodGitHub,
BotName: namePrefix,
GitHub: &types.ProvisionTokenSpecV2GitHub{
Allow: slices.Map(req.GitHub.Allow, func(allow machineIDWizardRequestGitHubAllow) *types.ProvisionTokenSpecV2GitHub_Rule {
return &types.ProvisionTokenSpecV2GitHub_Rule{
Repository: fmt.Sprintf("%s/%s", allow.Owner, allow.Repository),
RepositoryOwner: allow.Owner,
Workflow: allow.Workflow,
Environment: allow.Environment,
Actor: allow.Actor,
Ref: allow.Ref,
RefType: allow.RefType,
}
}),
EnterpriseServerHost: req.GitHub.EnterpriseServerHost,
EnterpriseSlug: req.GitHub.EnterpriseSlug,
StaticJWKS: req.GitHub.StaticJWKS,
},
},
}
roleCfg, err := tfgen.Generate(role, roleOpts...)
if err != nil {
return nil, trace.Wrap(err)
}
botCfg, err := tfgen.Generate(bot, botOpts...)
if err != nil {
return nil, trace.Wrap(err)
}
tokenCfg, err := tfgen.Generate(token, tfgen.WithResourceType("teleport_provision_token"))
if err != nil {
return nil, trace.Wrap(err)
}
var buf bytes.Buffer
err = templates.ExecuteTemplate(&buf, "machine-id-gha-k8s-wizard.tf.tmpl", struct {
RoleConfig string
BotConfig string
TokenConfig string
MajorVersion int64
ProxyAddr string
}{
RoleConfig: string(roleCfg),
BotConfig: string(botCfg),
TokenConfig: string(tokenCfg),
MajorVersion: teleport.SemVer().Major,
ProxyAddr: h.PublicProxyAddr(),
})
if err != nil {
return nil, trace.Wrap(err)
}
return machineIDGHAK8sWizardGenerateIaCResponse{Terraform: buf.String()}, nil
}
type machineIDWizardGenerateIaCRequest struct {
SourceType string `json:"source_type"`
DestinationType string `json:"destination_type"`
GitHub *machineIDWizardRequestGitHub `json:"github"`
Kubernetes *machineIDWizardRequestKubernetes `json:"kubernetes"`
}
type machineIDWizardRequestGitHub struct {
Allow []machineIDWizardRequestGitHubAllow `json:"allow"`
EnterpriseServerHost string `json:"enterprise_server_host"`
EnterpriseSlug string `json:"enterprise_slug"`
StaticJWKS string `json:"static_jwks"`
}
type machineIDWizardRequestGitHubAllow struct {
Repository string `json:"repository"`
Owner string `json:"owner"`
Workflow string `json:"workflow"`
Environment string `json:"environment"`
Actor string `json:"actor"`
Ref string `json:"ref"`
RefType string `json:"ref_type"`
}
type machineIDWizardRequestKubernetes struct {
Labels types.Labels `json:"labels"`
Groups []string `json:"groups"`
Users []string `json:"users"`
Resources []types.KubernetesResource `json:"resources"`
}
type machineIDGHAK8sWizardGenerateIaCResponse struct {
Terraform string `json:"terraform"`
}
var (
//go:embed templates/*
templateDir embed.FS
templates = template.Must(template.ParseFS(templateDir, "templates/*.tmpl"))
)
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/integrations/awsoidc"
"github.com/gravitational/teleport/lib/utils/oidc"
)
const (
// OIDCJWKWURI is the relative path where the OIDC IdP JWKS is located
OIDCJWKWURI = "/.well-known/jwks-oidc"
// OktaJWKSWellknownURI is the relative path where the Okta JWKS is located
OktaJWKSWellknownURI = "/.well-known/jwks-okta"
)
// openidConfiguration returns the openid-configuration for setting up the AWS OIDC Integration
func (h *Handler) openidConfiguration(_ http.ResponseWriter, _ *http.Request, _ httprouter.Params) (any, error) {
issuer, err := oidc.IssuerFromPublicAddress(h.cfg.PublicProxyAddr, "")
if err != nil {
return nil, trace.Wrap(err)
}
return oidc.OpenIDConfigurationForIssuer(issuer, issuer+OIDCJWKWURI), nil
}
// jwksOIDC returns all public keys used to sign JWT tokens for this cluster.
func (h *Handler) jwksOIDC(_ http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
return h.jwks(r.Context(), types.OIDCIdPCA, true)
}
// thumbprint returns the thumbprint as required by AWS when adding an OIDC Identity Provider.
// This is documented here:
// https://docs.aws.amazon.com/IAM/latest/UserGuide/id_roles_providers_create_oidc_verify-thumbprint.html
// Returns the thumbprint of the top intermediate CA that signed the TLS cert used to serve HTTPS requests.
// In case of a self signed certificate, then it returns the thumbprint of the TLS cert itself.
func (h *Handler) thumbprint(_ http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
return awsoidc.ThumbprintIdP(r.Context(), h.PublicProxyAddr(), h.cfg.InsecureMode)
}
// Teleport
// Copyright (C) 2024 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 (
"net/http"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types"
)
// jwksOkta returns public keys used to verify JWT tokens signed for use with Okta API Service App
// machine-to-machine authentication.
// https://developer.okta.com/docs/guides/implement-oauth-for-okta-serviceapp/main/
func (h *Handler) jwksOkta(_ http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
return h.jwks(r.Context(), types.OktaCA, false /* includeBlankKeyID */)
}
/*
* Teleport
* Copyright (C) 2026 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 "strconv"
// parseBoolWithDefault parses string for a boolean value,
// returns the default if the value is empty and
// returns error if the value is not recognized.
func parseBoolWithDefault(val string, def bool) (bool, error) {
if val == "" {
return def, nil
}
return strconv.ParseBool(val)
}
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/httplib"
)
// changePasswordReq is a request to change user password
type changePasswordReq struct {
// OldPassword is user current password
OldPassword []byte `json:"old_password"`
// NewPassword is user new password
NewPassword []byte `json:"new_password"`
// SecondFactorToken is user 2nd factor token
SecondFactorToken string `json:"second_factor_token"`
// WebauthnAssertionResponse is a Webauthn response
WebauthnAssertionResponse *wantypes.CredentialAssertionResponse `json:"webauthnAssertionResponse"`
}
// changePassword updates users password based on the old password.
func (h *Handler) changePassword(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext) (any, error) {
var req *changePasswordReq
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
protoReq := &proto.ChangePasswordRequest{
User: ctx.GetUser(),
OldPassword: req.OldPassword,
NewPassword: req.NewPassword,
SecondFactorToken: req.SecondFactorToken,
Webauthn: wantypes.CredentialAssertionResponseToProto(
req.WebauthnAssertionResponse,
),
}
if err := clt.ChangePassword(r.Context(), protoReq); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// Teleport
// Copyright (C) 2024 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"
"net"
"strconv"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/client/webclient"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/service/servicecfg"
)
// NetworkConfigGetter is a helper interface that allows to fetch the current proxy configuration.
type NetworkConfigGetter interface {
GetClusterNetworkingConfig(ctx context.Context) (types.ClusterNetworkingConfig, error)
}
// ProxySettings is a helper type that allows to fetch the current proxy configuration.
type ProxySettings struct {
// cfg is the Teleport service configuration.
ServiceConfig *servicecfg.Config
// proxySSHAddr is the address of the proxy ssh service. It can be assigned during runtime when a user set the
// proxy listener address to a random port (e.g. `127.0.0.1:0`).
ProxySSHAddr string
// accessPoint is the caching client connected to the auth server.
AccessPoint NetworkConfigGetter
}
// GetProxySettings allows returns current proxy configuration.
func (p *ProxySettings) GetProxySettings(ctx context.Context) (*webclient.ProxySettings, error) {
resp, err := p.AccessPoint.GetClusterNetworkingConfig(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
switch p.ServiceConfig.Version {
case defaults.TeleportConfigVersionV2, defaults.TeleportConfigVersionV3:
return p.buildProxySettingsV2(resp.GetProxyListenerMode(), resp.GetSSHDialTimeout()), nil
default:
return p.buildProxySettings(resp.GetProxyListenerMode(), resp.GetSSHDialTimeout()), nil
}
}
// buildProxySettings builds standard proxy configuration where proxy services are
// configured on different listeners. If the TLSRoutingEnabled flag is set and a proxy
// client support the TLSRouting dialer then the client will connect to the Teleport Proxy WebPort
// where incoming connections are routed to the proper proxy service based on TLS SNI ALPN routing information.
func (p *ProxySettings) buildProxySettings(proxyListenerMode types.ProxyListenerMode, sshDialTimeout time.Duration) *webclient.ProxySettings {
proxySettings := webclient.ProxySettings{
TLSRoutingEnabled: proxyListenerMode == types.ProxyListenerMode_Multiplex,
Kube: webclient.KubeProxySettings{
Enabled: p.ServiceConfig.Proxy.Kube.Enabled,
},
SSH: webclient.SSHProxySettings{
ListenAddr: p.ProxySSHAddr,
TunnelListenAddr: p.ServiceConfig.Proxy.ReverseTunnelListenAddr.String(),
WebListenAddr: p.ServiceConfig.Proxy.WebAddr.String(),
DialTimeout: sshDialTimeout,
},
ScopesEnabled: p.ServiceConfig.ScopesFeatures.Enabled,
GroupID: p.ServiceConfig.Proxy.ProxyGroupID,
}
p.setProxyPublicAddressesSettings(&proxySettings)
if !p.ServiceConfig.Proxy.MySQLAddr.IsEmpty() {
proxySettings.DB.MySQLListenAddr = p.ServiceConfig.Proxy.MySQLAddr.String()
}
if !p.ServiceConfig.Proxy.PostgresAddr.IsEmpty() {
proxySettings.DB.PostgresListenAddr = p.ServiceConfig.Proxy.PostgresAddr.String()
}
if !p.ServiceConfig.Proxy.MongoAddr.IsEmpty() {
proxySettings.DB.MongoListenAddr = p.ServiceConfig.Proxy.MongoAddr.String()
}
if p.ServiceConfig.Proxy.Kube.Enabled {
proxySettings.Kube.ListenAddr = p.ServiceConfig.Proxy.Kube.ListenAddr.String()
}
return &proxySettings
}
// buildProxySettingsV2 builds the v2 proxy settings where teleport proxies can start only on a single listener.
func (p *ProxySettings) buildProxySettingsV2(proxyListenerMode types.ProxyListenerMode, sshDialTimeout time.Duration) *webclient.ProxySettings {
multiplexAddr := p.ServiceConfig.Proxy.WebAddr.String()
settings := p.buildProxySettings(proxyListenerMode, sshDialTimeout)
if proxyListenerMode == types.ProxyListenerMode_Multiplex {
settings.SSH.ListenAddr = multiplexAddr
settings.SSH.TunnelListenAddr = multiplexAddr
settings.SSH.WebListenAddr = multiplexAddr
settings.Kube.ListenAddr = multiplexAddr
settings.DB.MySQLListenAddr = multiplexAddr
settings.DB.PostgresListenAddr = multiplexAddr
}
return settings
}
func (p *ProxySettings) setProxyPublicAddressesSettings(settings *webclient.ProxySettings) {
if len(p.ServiceConfig.Proxy.PublicAddrs) > 0 {
settings.SSH.PublicAddr = p.ServiceConfig.Proxy.PublicAddrs[0].String()
}
if len(p.ServiceConfig.Proxy.SSHPublicAddrs) > 0 {
settings.SSH.SSHPublicAddr = p.ServiceConfig.Proxy.SSHPublicAddrs[0].String()
}
if len(p.ServiceConfig.Proxy.TunnelPublicAddrs) > 0 {
settings.SSH.TunnelPublicAddr = p.ServiceConfig.Proxy.TunnelPublicAddrs[0].String()
}
if len(p.ServiceConfig.Proxy.Kube.PublicAddrs) > 0 {
settings.Kube.PublicAddr = p.ServiceConfig.Proxy.Kube.PublicAddrs[0].String()
}
if len(p.ServiceConfig.Proxy.MySQLPublicAddrs) > 0 {
settings.DB.MySQLPublicAddr = p.ServiceConfig.Proxy.MySQLPublicAddrs[0].String()
}
if len(p.ServiceConfig.Proxy.MongoPublicAddrs) > 0 {
settings.DB.MongoPublicAddr = p.ServiceConfig.Proxy.MongoPublicAddrs[0].String()
}
settings.DB.PostgresPublicAddr = p.getPostgresPublicAddr()
}
// getPostgresPublicAddr returns the proxy PostgresPublicAddrs based on whether the Postgres proxy service
// was configured on separate listener. For backward compatibility if PostgresPublicAddrs was not provided.
// Proxy will reuse the PostgresPublicAddrs field to propagate postgres service address to legacy tsh clients.
func (p *ProxySettings) getPostgresPublicAddr() string {
if len(p.ServiceConfig.Proxy.PostgresPublicAddrs) > 0 {
return p.ServiceConfig.Proxy.PostgresPublicAddrs[0].String()
}
if p.ServiceConfig.Proxy.PostgresAddr.IsEmpty() {
return ""
}
// DELETE IN 9.0.0
// If the PostgresPublicAddrs address was not set propagate separate postgres service listener address
// to legacy tsh clients reusing PostgresPublicAddrs field.
var host string
if len(p.ServiceConfig.Proxy.PublicAddrs) > 0 {
// Get proxy host address from public address.
host = p.ServiceConfig.Proxy.PublicAddrs[0].Host()
} else {
host = p.ServiceConfig.Proxy.WebAddr.Host()
}
return net.JoinHostPort(host, strconv.Itoa(p.ServiceConfig.Proxy.PostgresAddr.Port(defaults.PostgresListenPort)))
}
/*
* Teleport
* Copyright (C) 2025 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 (
"errors"
"fmt"
"io"
"net/http"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"google.golang.org/protobuf/types/known/durationpb"
recordingmetadatav1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/recordingmetadata/v1"
"github.com/gravitational/teleport/lib/reversetunnelclient"
)
type sessionRecordingMessageType string
const (
recordingThumbnailMessageType sessionRecordingMessageType = "thumbnail"
recordingMetadataMessageType sessionRecordingMessageType = "metadata"
recordingErrorMessageType sessionRecordingMessageType = "error"
)
type sessionRecordingErrorResponse struct {
Error string `json:"error"`
}
// sessionRecordingMessageWrapper is a wrapper for session recording messages sent over WebSocket.
// This makes it easier to have strongly typed messages on the frontend, switching on the `Type` field.
type sessionRecordingMessageWrapper struct {
Type sessionRecordingMessageType `json:"type"`
Data any `json:"data"`
}
// getSessionRecordingMetadata handles the WebSocket connection to stream session recording metadata and thumbnails.
// The metadata is loaded over a websocket connection to avoid gRPC message size limits.
// It sends metadata and thumbnails as JSON messages to the client.
func (h *Handler) getSessionRecordingMetadata(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (interface{}, error) {
sessionID := p.ByName("session_id")
if sessionID == "" {
return nil, trace.BadParameter("missing session ID in request URL")
}
ctx := r.Context()
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
sendMessage(ws, recordingErrorMessageType, sessionRecordingErrorResponse{
Error: err.Error(),
})
return nil, nil
}
stream, err := clt.RecordingMetadataServiceClient().GetMetadata(ctx, recordingmetadatav1.GetMetadataRequest_builder{
SessionId: sessionID,
}.Build())
if err != nil {
sendMessage(ws, recordingErrorMessageType, sessionRecordingErrorResponse{
Error: err.Error(),
})
return nil, nil
}
for {
chunk, err := stream.Recv()
if err != nil {
if trace.IsNotFound(err) {
sendMessage(ws, recordingErrorMessageType, sessionRecordingErrorResponse{
Error: fmt.Sprintf("metadata for session %q not found", sessionID),
})
return nil, nil
}
if errors.Is(err, io.EOF) {
break
}
h.logger.ErrorContext(ctx, "failed to receive chunk", "session_id", sessionID, "error", err)
sendMessage(ws, recordingErrorMessageType, sessionRecordingErrorResponse{
Error: err.Error(),
})
return nil, nil
}
if chunk.GetMetadata() != nil {
if err := sendMessage(ws, recordingMetadataMessageType, encodeSessionRecordingMetadata(chunk.GetMetadata())); err != nil {
h.logger.ErrorContext(ctx, "failed to send metadata", "session_id", sessionID, "error", err)
return nil, nil
}
continue
}
if chunk.GetFrame() != nil {
if err := sendMessage(ws, recordingThumbnailMessageType, EncodeSessionRecordingThumbnail(chunk.GetFrame())); err != nil {
h.logger.ErrorContext(ctx, "failed to send thumbnail", "session_id", sessionID, "error", err)
return nil, nil
}
continue
}
h.logger.ErrorContext(ctx, "received nil frame in metadata stream")
sendMessage(ws, recordingErrorMessageType, sessionRecordingErrorResponse{
Error: trace.BadParameter("received nil frame").Error(),
})
return nil, nil
}
return nil, nil
}
func sendMessage(ws *websocket.Conn, msgType sessionRecordingMessageType, data interface{}) error {
return ws.WriteJSON(sessionRecordingMessageWrapper{
Type: msgType,
Data: data,
})
}
func (h *Handler) getSessionRecordingThumbnail(
w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster,
) (any, error) {
sessionId := p.ByName("session_id")
if sessionId == "" {
return nil, trace.BadParameter("session_id is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
response, err := clt.RecordingMetadataServiceClient().GetThumbnail(r.Context(), recordingmetadatav1.GetThumbnailRequest_builder{
SessionId: sessionId,
}.Build())
if err != nil {
if trace.IsNotFound(err) {
return nil, trace.NotFound("thumbnail not found for session %q", sessionId)
}
return nil, trace.Wrap(err)
}
if !response.HasThumbnail() {
return nil, trace.NotFound("thumbnail not found for session %q", sessionId)
}
return EncodeSessionRecordingThumbnail(response.GetThumbnail()), nil
}
type baseEvent struct {
StartOffset int64 `json:"startTime"`
EndOffset int64 `json:"endTime"`
Type string `json:"type"`
}
type resizeEvent struct {
baseEvent
Cols int32 `json:"cols"`
Rows int32 `json:"rows"`
}
func (resizeEvent) isSessionRecordingEvent() {}
type joinEvent struct {
baseEvent
User string `json:"user"`
}
func (joinEvent) isSessionRecordingEvent() {}
type inactivityEvent struct {
baseEvent
}
func (inactivityEvent) isSessionRecordingEvent() {}
type sessionRecordingEvent interface {
isSessionRecordingEvent()
}
type sessionRecordingMetadata struct {
Duration int64 `json:"duration"`
Events []sessionRecordingEvent `json:"events"`
StartCols int32 `json:"startCols"`
StartRows int32 `json:"startRows"`
StartTime int64 `json:"startTime"`
EndTime int64 `json:"endTime"`
ClusterName string `json:"clusterName"`
ResourceName string `json:"resourceName,omitempty"`
User string `json:"user,omitempty"`
Type string `json:"type,omitempty"`
}
func pbTypeToString(t recordingmetadatav1.SessionRecordingType) string {
switch t {
case recordingmetadatav1.SessionRecordingType_SESSION_RECORDING_TYPE_SSH:
return "ssh"
case recordingmetadatav1.SessionRecordingType_SESSION_RECORDING_TYPE_KUBERNETES:
return "k8s"
case recordingmetadatav1.SessionRecordingType_SESSION_RECORDING_TYPE_WINDOWS_DESKTOP:
return "desktop"
default:
return "unknown"
}
}
// encodeSessionRecordingMetadata converts the session recording metadata to a format more suitable for the frontend
// to use.
func encodeSessionRecordingMetadata(metadata *recordingmetadatav1.SessionRecordingMetadata) sessionRecordingMetadata {
result := sessionRecordingMetadata{
Duration: convertDurationToMs(metadata.GetDuration()),
StartCols: metadata.GetStartCols(),
StartRows: metadata.GetStartRows(),
Events: make([]sessionRecordingEvent, 0, len(metadata.GetEvents())),
StartTime: metadata.GetStartTime().AsTime().Unix(),
EndTime: metadata.GetEndTime().AsTime().Unix(),
ClusterName: metadata.GetClusterName(),
ResourceName: metadata.GetResourceName(),
User: metadata.GetUser(),
Type: pbTypeToString(metadata.GetType()),
}
for _, event := range metadata.GetEvents() {
base := baseEvent{
StartOffset: convertDurationToMs(event.GetStartOffset()),
EndOffset: convertDurationToMs(event.GetEndOffset()),
}
switch event.WhichEvent() {
case recordingmetadatav1.SessionRecordingEvent_Inactivity_case:
base.Type = "inactivity"
result.Events = append(result.Events, inactivityEvent{baseEvent: base})
case recordingmetadatav1.SessionRecordingEvent_Join_case:
base.Type = "join"
result.Events = append(result.Events, joinEvent{
baseEvent: base,
User: event.GetJoin().GetUser(),
})
case recordingmetadatav1.SessionRecordingEvent_Resize_case:
base.Type = "resize"
result.Events = append(result.Events, resizeEvent{
baseEvent: base,
Cols: event.GetResize().GetCols(),
Rows: event.GetResize().GetRows(),
})
}
}
return result
}
// SessionRecordingThumbnailResponse is the web API representation of a session recording thumbnail.
type SessionRecordingThumbnailResponse struct {
Svg string `json:"svg"`
Cols int32 `json:"cols"`
Rows int32 `json:"rows"`
CursorX int32 `json:"cursorX"`
CursorY int32 `json:"cursorY"`
CursorVisible bool `json:"cursorVisible"`
StartOffset int64 `json:"startOffset"`
EndOffset int64 `json:"endOffset"`
Png []byte `json:"png,omitempty"`
ScreenWidth int32 `json:"screenWidth,omitempty"`
ScreenHeight int32 `json:"screenHeight,omitempty"`
}
// EncodeSessionRecordingThumbnail converts the session recording thumbnail to a format more suitable for the frontend.
func EncodeSessionRecordingThumbnail(thumbnail *recordingmetadatav1.SessionRecordingThumbnail) SessionRecordingThumbnailResponse {
return SessionRecordingThumbnailResponse{
Svg: string(thumbnail.GetSvg()),
Cols: thumbnail.GetCols(),
Rows: thumbnail.GetRows(),
CursorX: thumbnail.GetCursorX(),
CursorY: thumbnail.GetCursorY(),
CursorVisible: thumbnail.GetCursorVisible(),
StartOffset: convertDurationToMs(thumbnail.GetStartOffset()),
EndOffset: convertDurationToMs(thumbnail.GetEndOffset()),
Png: thumbnail.GetPng(),
ScreenWidth: thumbnail.GetScreenWidth(),
ScreenHeight: thumbnail.GetScreenHeight(),
}
}
func convertDurationToMs(d *durationpb.Duration) int64 {
if d == nil {
return 0
}
return d.Seconds*1000 + int64(d.Nanos/1000000)
}
/**
* Teleport
* Copyright (C) 2025 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"
"encoding/binary"
"fmt"
"log/slog"
"net/http"
"runtime/debug"
"sync"
"time"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/hinshun/vt10x"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/metadata"
apievents "github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/events"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/utils"
)
// maxRequestRange is the maximum allowed time range for a request
const maxRequestRange = 10 * time.Minute
const websocketCloseTimeout = 5 * time.Second
// websocketMessage represents a message to be written to the websocket
type websocketMessage struct {
messageType int
data []byte
}
type recordingTerminal struct {
sync.Mutex
vt vt10x.Terminal
}
type recordingStream struct {
sync.Mutex
eventsChan <-chan apievents.AuditEvent
errorsChan <-chan error
lastEndTime time.Duration
bufferedEvent apievents.AuditEvent
}
// recordingPlayback manages session event streaming
type recordingPlayback struct {
ctx context.Context
cancel context.CancelFunc
clt events.SessionStreamer
sessionID string
logger *slog.Logger
mu sync.Mutex
cancelActiveTask context.CancelFunc
wg sync.WaitGroup
ws *websocket.Conn
writeChan chan websocketMessage
closeSent bool // tracks if a close message has been sent
stream recordingStream
terminal recordingTerminal
}
// fetchRequest represents a request for session events.
type fetchRequest struct {
requestType requestType
startOffset time.Duration
endOffset time.Duration
requestID int
requestCurrentScreen bool
}
// sessionEvent represents a single session event with its type, timestamp, and data.
type sessionEvent struct {
eventType responseType
timeOffset time.Duration
data []byte
}
func (h *Handler) recordingPlaybackWS(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
ws *websocket.Conn,
) (interface{}, error) {
sessionID := p.ByName("session_id")
if sessionID == "" {
h.closeWebsocketWithError(r.Context(), ws, sessionID, trace.BadParameter("missing session ID in request URL"))
return nil, nil
}
ctx := r.Context()
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
h.closeWebsocketWithError(ctx, ws, sessionID, trace.Wrap(err, "failed to get user client"))
return nil, nil
}
playback := newRecordingPlayback(ctx, ws, clt, sessionID, h.logger)
playback.run()
return nil, nil
}
func (h *Handler) closeWebsocketWithError(ctx context.Context, ws *websocket.Conn, sessionID string, err error) {
data := []byte(err.Error())
totalSize := responseHeaderSize + len(data)
buf := make([]byte, totalSize)
encodeEvent(buf, 0, eventTypeError, 0, data, 0)
if err := ws.WriteMessage(websocket.BinaryMessage, buf); err != nil {
h.logger.ErrorContext(ctx, "failed to send event", "session_id", sessionID, "error", err)
}
deadline := time.Now().Add(websocketCloseTimeout)
// Send close frame to initiate graceful shutdown
if err := ws.SetWriteDeadline(deadline); err != nil {
h.logger.DebugContext(ctx, "failed to set write deadline", "session_id", sessionID, "error", err)
}
if err := ws.WriteMessage(websocket.CloseMessage, websocket.FormatCloseMessage(websocket.CloseNormalClosure, "")); err != nil {
h.logger.DebugContext(ctx, "failed to send close message", "session_id", sessionID, "error", err)
}
// Wait for peer's close frame response (or timeout)
if err := ws.SetReadDeadline(deadline); err != nil {
h.logger.DebugContext(ctx, "failed to set read deadline", "session_id", sessionID, "error", err)
}
// Log if we got something other than a close acknowledgement
if _, _, err := ws.ReadMessage(); err != nil && !websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) {
h.logger.DebugContext(ctx, "received non-close message while waiting for close acknowledgement",
"session_id", sessionID, "error", err)
}
// Finally close the underlying connection
ws.Close()
}
// newRecordingPlayback creates a new session recording playback handler.
// This provides a way for the client to request session events within a specific time range, as well as the current
// terminal screen state at a given time (when seeking).
// This allows for faster seeking without having to send the client extra events to reconstruct the terminal state.
func newRecordingPlayback(ctx context.Context, ws *websocket.Conn, clt events.SessionStreamer, sessionID string, logger *slog.Logger) *recordingPlayback {
ctx, cancel := context.WithCancel(ctx)
s := &recordingPlayback{
ctx: ctx,
cancel: cancel,
clt: clt,
sessionID: sessionID,
logger: logger,
ws: ws,
writeChan: make(chan websocketMessage),
}
return s
}
// run starts the recording playback handler.
//
// A recovered panic on the read-loop goroutine (for example, vt10x tripping
// over a corrupt recording during handleFetchRequest) is logged rather than
// propagating and crashing the proxy. The streamEvents goroutine has its own
// defer/recover since it runs independently.
func (s *recordingPlayback) run() {
defer s.cleanup()
defer func() {
if r := recover(); r != nil {
s.logger.ErrorContext(s.ctx, "panic on recording playback read loop",
"session_id", s.sessionID,
"panic", r,
"stack", string(debug.Stack()),
)
}
}()
go s.writeLoop()
s.readLoop()
}
// cleanup cleans up the recording playback resources.
func (s *recordingPlayback) cleanup() {
s.cancel()
s.mu.Lock()
// Only send close message if we haven't already sent one
if !s.closeSent {
select {
case s.writeChan <- websocketMessage{
messageType: websocket.CloseMessage,
data: websocket.FormatCloseMessage(websocket.CloseNormalClosure, ""),
}:
case <-time.After(websocketCloseTimeout):
}
}
s.mu.Unlock()
// Wait for any active task to complete
s.wg.Wait()
close(s.writeChan)
// Wait for peer's close frame response (or timeout)
deadline := time.Now().Add(websocketCloseTimeout)
if err := s.ws.SetReadDeadline(deadline); err != nil {
s.logger.DebugContext(s.ctx, "failed to set read deadline", "session_id", s.sessionID, "error", err)
}
// Log if we got something other than a close acknowledgement
if _, _, err := s.ws.ReadMessage(); err != nil && !websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) {
s.logger.DebugContext(s.ctx, "received non-close message while waiting for close acknowledgement",
"session_id", s.sessionID, "error", err)
}
// Finally close the underlying connection
s.ws.Close()
}
// writeLoop handles all websocket writes from a dedicated goroutine.
func (s *recordingPlayback) writeLoop() {
for {
select {
case <-s.ctx.Done():
return
case msg, ok := <-s.writeChan:
if !ok {
return
}
if err := s.ws.SetWriteDeadline(time.Now().Add(10 * time.Second)); err != nil {
s.logWebsocketError(err)
return
}
if err := s.ws.WriteMessage(msg.messageType, msg.data); err != nil {
s.logWebsocketError(err)
return
}
// If we just sent a close message, exit the loop
if msg.messageType == websocket.CloseMessage {
// Mark that we're sending a close message
s.mu.Lock()
s.closeSent = true
s.mu.Unlock()
return
}
}
}
}
// logWebsocketError handles errors that occur during websocket writes.
func (s *recordingPlayback) logWebsocketError(err error) {
if !websocket.IsCloseError(err, websocket.CloseNormalClosure, websocket.CloseGoingAway) &&
!utils.IsOKNetworkError(err) {
s.logger.ErrorContext(s.ctx, "websocket write error", "error", err)
}
}
// readLoop reads messages from the websocket connection and processes them.
func (s *recordingPlayback) readLoop() {
for {
msgType, data, err := s.ws.ReadMessage()
if err != nil {
s.logWebsocketError(err)
return
}
if msgType != websocket.BinaryMessage {
s.logger.ErrorContext(s.ctx, "ignoring non-binary websocket message", "session_id", s.sessionID, "type", msgType)
// Send close message through the write channel
select {
case s.writeChan <- websocketMessage{
messageType: websocket.CloseMessage,
data: websocket.FormatCloseMessage(websocket.CloseUnsupportedData, "only binary messages are supported"),
}:
case <-time.After(1 * time.Second):
s.logger.ErrorContext(s.ctx, "timeout sending close message", "session_id", s.sessionID)
}
return
}
req, err := decodeBinaryRequest(data)
if err != nil {
s.logger.WarnContext(s.ctx, "failed to decode request", "session_id", s.sessionID, "error", err)
continue
}
switch req.requestType {
case requestTypeFetch:
s.handleFetchRequest(req)
default:
s.sendError(trace.BadParameter("unknown request type: %d", req.requestType), req.requestID)
s.logger.ErrorContext(s.ctx, "received unknown request type", "session_id", s.sessionID, "type", req.requestType)
}
}
}
// createTaskContext creates a new context for a task and cancels any previous task.
// A task context is used to manage the lifecycle of a fetch request, ensuring that only one fetch request is active at a time.
func (s *recordingPlayback) createTaskContext() context.Context {
s.mu.Lock()
if s.cancelActiveTask != nil {
// Cancel the active task first
s.cancelActiveTask()
s.mu.Unlock()
// Wait for streamEvents to terminate before continuing
// We unlock the mutex while waiting to avoid deadlock
s.wg.Wait()
s.mu.Lock()
}
ctx, taskCancel := context.WithCancel(s.ctx)
s.cancelActiveTask = taskCancel
s.mu.Unlock()
return ctx
}
// handleFetchRequest processes a fetch request for session events.
func (s *recordingPlayback) handleFetchRequest(req *fetchRequest) {
if err := validateRequest(req); err != nil {
s.sendError(err, req.requestID)
return
}
ctx := s.createTaskContext()
s.stream.Lock()
// start the stream if it doesn't exist or if we need to go back in time
needNewStream := s.stream.eventsChan == nil || req.startOffset < s.stream.lastEndTime
if needNewStream {
events, errors := s.clt.StreamSessionEvents(
metadata.WithSessionRecordingFormatContext(s.ctx, teleport.PTY),
session.ID(s.sessionID),
0,
)
if events == nil || errors == nil {
s.sendError(fmt.Errorf("failed to start session event stream"), req.requestID)
s.stream.Unlock()
return
}
s.stream.eventsChan = events
s.stream.errorsChan = errors
s.stream.lastEndTime = 0
s.terminal.Lock()
s.terminal.vt = vt10x.New()
s.terminal.Unlock()
}
s.stream.lastEndTime = req.endOffset
eventsChan := s.stream.eventsChan
errorsChan := s.stream.errorsChan
s.stream.Unlock()
s.wg.Add(1)
go func() {
defer s.wg.Done()
s.streamEvents(ctx, req, eventsChan, errorsChan)
}()
}
// streamEvents streams session events to the client.
//
// A recovered panic (e.g. vt10x tripping over a corrupt recording) is logged
// and reported to the client as an error followed by a stop event, rather
// than crashing the proxy. The stop event is required: the web client only
// clears its loading state on stop, so skipping it would leave playback
// stuck in loading after a malformed recording.
func (s *recordingPlayback) streamEvents(ctx context.Context, req *fetchRequest, eventsChan <-chan apievents.AuditEvent, errorsChan <-chan error) {
startSent := false
screenSent := false
inTimeRange := false
const maxBatchSize = 200
eventBatch := make([]sessionEvent, 0, maxBatchSize)
flushBatch := func() {
// Send start event if not already sent
if !startSent {
s.sendEvent(eventTypeStart, req.startOffset, nil, req.requestID)
startSent = true
}
if len(eventBatch) == 0 {
return
}
s.sendEventBatch(eventBatch, req.requestID)
eventBatch = eventBatch[:0]
}
addToBatch := func(eventType responseType, timeOffset time.Duration, data []byte) {
eventBatch = append(eventBatch, sessionEvent{eventType, timeOffset, data})
if len(eventBatch) >= maxBatchSize {
flushBatch()
}
}
sendStop := func() {
// Send start event if not already sent
if !startSent {
s.sendEvent(eventTypeStart, req.startOffset, nil, req.requestID)
startSent = true
}
s.sendEvent(eventTypeStop, 0, encodeTime(req.startOffset, req.endOffset), req.requestID)
}
defer func() {
if r := recover(); r != nil {
s.logger.ErrorContext(s.ctx, "panic while streaming session recording events",
"session_id", s.sessionID,
"panic", r,
"stack", string(debug.Stack()),
)
// cleanup() cancels s.ctx before closing s.writeChan, so if the context is already done we must not touch
// writeChan — a concurrent close would turn into a panic.
if s.ctx.Err() != nil {
return
}
s.sendError(trace.Errorf("internal error while streaming session recording"), req.requestID)
sendStop()
}
}()
// process an event, returning a boolean indicating if the events should continue being
// processed (i.e. returns false once we have reached the end of the requested time window)
processEvent := func(evt apievents.AuditEvent) bool {
eventTime := getEventTime(evt)
inTimeRange = eventTime >= req.startOffset && eventTime <= req.endOffset
if inTimeRange && req.requestCurrentScreen && !screenSent {
flushBatch()
s.sendCurrentScreen(req.requestID, eventTime)
screenSent = true
}
if eventTime > req.endOffset {
s.stream.Lock()
// store the event for the next request as it is outside the current time range
// and won't be returned by the stream on the next request
// this will only store print or end events as they are the only ones with a timestamp
s.stream.bufferedEvent = evt
s.stream.Unlock()
return false
}
switch evt := evt.(type) {
case *apievents.SessionStart:
if err := s.resizeTerminal(evt.TerminalSize); err != nil {
s.logger.ErrorContext(s.ctx, "failed to resize terminal", "session_id", s.sessionID, "error", err)
// continue returning events even if resize fails
}
if inTimeRange {
addToBatch(eventTypeSessionStart, 0, []byte(evt.TerminalSize))
}
case *apievents.SessionPrint:
// defer Unlock so a panic in vt.Write (caught by streamEvents' defer/recover) can't leave s.terminal locked and
// wedge later playback requests on the same websocket.
func() {
s.terminal.Lock()
defer s.terminal.Unlock()
if _, err := s.terminal.vt.Write(evt.Data); err != nil {
s.logger.ErrorContext(s.ctx, "failed to write to terminal", "session_id", s.sessionID, "error", err)
}
}()
if inTimeRange {
addToBatch(eventTypeSessionPrint, eventTime, evt.Data)
}
case *apievents.SessionEnd:
endTime := evt.EndTime.Sub(evt.StartTime)
if inTimeRange {
addToBatch(eventTypeSessionEnd, endTime, []byte(evt.EndTime.Format(time.RFC3339)))
}
return false
case *apievents.Resize:
if err := s.resizeTerminal(evt.TerminalSize); err != nil {
s.logger.ErrorContext(s.ctx, "failed to resize terminal", "session_id", s.sessionID, "error", err)
// continue returning events even if resize fails
}
// always add resize events as they do not have a timestamp
addToBatch(eventTypeResize, 0, []byte(evt.TerminalSize))
}
return true
}
s.stream.Lock()
buffered := s.stream.bufferedEvent
s.stream.bufferedEvent = nil
s.stream.Unlock()
if buffered != nil {
// process any buffered event from a previous request first
// the processEvent will ignore it if it's outside the requested time range
_ = processEvent(buffered)
}
for {
select {
case <-ctx.Done():
// Don't send any more events after context cancellation
return
case err := <-errorsChan:
flushBatch()
if err != nil {
s.sendError(err, req.requestID)
}
sendStop()
return
case evt, ok := <-eventsChan:
if !ok {
flushBatch()
// Send screen if requested and not already sent
// This handles the case where the stream ends, but we haven't sent the screen yet
if req.requestCurrentScreen && !screenSent {
s.sendCurrentScreen(req.requestID, req.startOffset)
}
sendStop()
return
}
if !processEvent(evt) {
flushBatch()
// Send screen if requested and not already sent when we reach the end of the time range
// (i.e. there was no event in the time range)
if req.requestCurrentScreen && !screenSent {
s.sendCurrentScreen(req.requestID, req.startOffset)
}
sendStop()
return
}
}
}
}
// resizeTerminal resizes the terminal based on the provided size string.
func (s *recordingPlayback) resizeTerminal(size string) error {
params, err := session.UnmarshalTerminalParams(size)
if err != nil {
return trace.Wrap(err)
}
s.terminal.Lock()
defer s.terminal.Unlock()
s.terminal.vt.Resize(params.W, params.H)
return nil
}
// writeMessage sends a message through the write channel.
func (s *recordingPlayback) writeMessage(data []byte) error {
select {
case <-s.ctx.Done():
return s.ctx.Err()
case s.writeChan <- websocketMessage{messageType: websocket.BinaryMessage, data: data}:
return nil
case <-time.After(10 * time.Second):
return fmt.Errorf("timeout sending message")
}
}
// sendEvent sends a single event to the client.
func (s *recordingPlayback) sendEvent(eventType responseType, timeOffset time.Duration, data []byte, requestID int) {
totalSize := responseHeaderSize + len(data)
buf := make([]byte, totalSize)
encodeEvent(buf, 0, eventType, timeOffset, data, requestID)
if err := s.writeMessage(buf); err != nil {
s.logger.ErrorContext(s.ctx, "failed to send event", "session_id", s.sessionID, "error", err)
}
}
// sendEventBatch sends a batch of events to the client.
func (s *recordingPlayback) sendEventBatch(batch []sessionEvent, requestID int) {
totalSize := responseHeaderSize
for _, evt := range batch {
totalSize += responseHeaderSize + len(evt.data)
}
buf := make([]byte, totalSize)
buf[0] = byte(eventTypeBatch)
binary.BigEndian.PutUint32(buf[1:5], uint32(len(batch)))
binary.BigEndian.PutUint32(buf[5:9], uint32(requestID))
offset := responseHeaderSize
for _, evt := range batch {
encodeEvent(buf, offset, evt.eventType, evt.timeOffset, evt.data, requestID)
offset += responseHeaderSize + len(evt.data)
}
if err := s.writeMessage(buf); err != nil {
s.logger.ErrorContext(s.ctx, "failed to send event batch",
"session_id", s.sessionID,
"error", err,
"batch_size", len(batch),
"buffer_size", totalSize)
}
}
// sendError sends an error event to the client.
func (s *recordingPlayback) sendError(err error, requestID int) {
if trace.IsAccessDenied(err) {
s.sendEvent(eventTypeError, 0, []byte("Session recording not found"), requestID)
return
}
s.sendEvent(eventTypeError, 0, []byte(err.Error()), requestID)
}
// sendCurrentScreen sends the current terminal screen state to the client.
func (s *recordingPlayback) sendCurrentScreen(requestID int, timeOffset time.Duration) {
// defer Unlock so a panic in any vt10x call (caught by streamEvents' defer/recover) can't leave s.terminal locked.
state, cols, rows, cursor := func() (vt10x.TerminalState, int, int, vt10x.Cursor) {
s.terminal.Lock()
defer s.terminal.Unlock()
state := s.terminal.vt.DumpState()
cols, rows := s.terminal.vt.Size()
cursor := s.terminal.vt.Cursor()
return state, cols, rows, cursor
}()
data := encodeScreenEvent(state, cols, rows, cursor)
s.sendEvent(eventTypeScreen, timeOffset, data, requestID)
}
// getEventTime extracts the event time from an AuditEvent.
func getEventTime(evt apievents.AuditEvent) time.Duration {
switch evt := evt.(type) {
case *apievents.SessionPrint:
return time.Duration(evt.DelayMilliseconds) * time.Millisecond
case *apievents.SessionEnd:
return evt.EndTime.Sub(evt.StartTime)
default:
return 0
}
}
/**
* Teleport
* Copyright (C) 2025 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 (
"bytes"
"encoding/binary"
"fmt"
"time"
"github.com/gravitational/trace"
"github.com/hinshun/vt10x"
"github.com/gravitational/teleport/lib/terminal"
)
const (
// requestHeaderSize is the size of the request header (event type, start time, end time, request ID, and current screen flag)
requestHeaderSize = 22
// responseHeaderSize is the size of the response header (event type, timestamp, data size, and request ID)
responseHeaderSize = 17
)
type requestType byte
// Identifies requests coming from the client (web UI)
const (
// requestTypeFetch requests event data
requestTypeFetch requestType = 1
)
type responseType byte
// Response types sent back to the client
const (
// eventTypeStart indicates the start of a response of events
eventTypeStart responseType = 1
// eventTypeStop indicates the stop of a response of events
eventTypeStop responseType = 2
// eventTypeError indicates an error
eventTypeError responseType = 3
// eventTypeSessionStart indicates session started
eventTypeSessionStart responseType = 4
// eventTypeSessionPrint contains terminal output
eventTypeSessionPrint responseType = 5
// eventTypeSessionEnd indicates session ended
eventTypeSessionEnd responseType = 6
// eventTypeResize indicates terminal resize
eventTypeResize responseType = 7
// eventTypeScreen contains terminal screen state
eventTypeScreen responseType = 8
// eventTypeBatch indicates a batch of events
eventTypeBatch responseType = 9
)
// encodeScreenEvent encodes the current terminal screen state into a byte slice.
func encodeScreenEvent(state vt10x.TerminalState, cols, rows int, cursor vt10x.Cursor) []byte {
var buf bytes.Buffer
buf.Write(make([]byte, responseHeaderSize))
terminal.VtStateToANSI(&buf, state)
eventData := buf.Bytes()
eventData[0] = byte(eventTypeScreen)
binary.BigEndian.PutUint32(eventData[1:5], uint32(cols))
binary.BigEndian.PutUint32(eventData[5:9], uint32(rows))
binary.BigEndian.PutUint32(eventData[9:13], uint32(cursor.X))
binary.BigEndian.PutUint32(eventData[13:17], uint32(cursor.Y))
binary.BigEndian.PutUint32(eventData[17:21], uint32(len(eventData)-responseHeaderSize))
return eventData
}
// encodeEvent encodes a session event into a byte slice.
func encodeEvent(buf []byte, offset int, eventType responseType, timeOffset time.Duration, data []byte, requestID int) {
buf[offset] = byte(eventType)
binary.BigEndian.PutUint64(buf[offset+1:offset+9], uint64(timeOffset/time.Millisecond))
binary.BigEndian.PutUint32(buf[offset+9:offset+13], uint32(len(data)))
binary.BigEndian.PutUint32(buf[offset+13:offset+17], uint32(requestID))
copy(buf[offset+responseHeaderSize:], data)
}
// encodeTime encodes the start and end times into a byte slice.
func encodeTime(startTime, endTime time.Duration) []byte {
buf := make([]byte, 16)
binary.BigEndian.PutUint64(buf, uint64(startTime/time.Millisecond))
binary.BigEndian.PutUint64(buf[8:], uint64(endTime/time.Millisecond))
return buf
}
// decodeBinaryRequest decodes a binary request from the client.
func decodeBinaryRequest(data []byte) (*fetchRequest, error) {
if len(data) != requestHeaderSize {
return nil, trace.BadParameter("invalid request size: expected %d bytes, got %d bytes", requestHeaderSize, len(data))
}
req := &fetchRequest{
requestType: requestType(data[0]),
startOffset: time.Duration(binary.BigEndian.Uint64(data[1:9])) * time.Millisecond,
endOffset: time.Duration(binary.BigEndian.Uint64(data[9:17])) * time.Millisecond,
requestID: int(binary.BigEndian.Uint32(data[17:21])),
requestCurrentScreen: data[21] == 1,
}
return req, nil
}
// validateRequest validates the fetch request parameters.
func validateRequest(req *fetchRequest) error {
if req.startOffset < 0 || req.endOffset < 0 || req.endOffset < req.startOffset {
return fmt.Errorf("invalid time range (%v, %v)", req.startOffset, req.endOffset)
}
if req.endOffset-req.startOffset > maxRequestRange {
return trace.LimitExceeded("time range too large, max is %s", maxRequestRange)
}
return nil
}
/*
* 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"
"iter"
"net/http"
"net/url"
"strings"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
kyaml "k8s.io/apimachinery/pkg/util/yaml"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/api/constants"
kubeproto "github.com/gravitational/teleport/api/gen/proto/go/teleport/kube/v1"
"github.com/gravitational/teleport/api/mfa"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/lib/auth"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/web/ui"
)
func (h *Handler) listRolesHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
return listRoles(clt, values)
}
func listRoles(clt resourcesAPIGetter, values url.Values) (*listResourcesWithoutCountGetResponse, error) {
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
skipSystemRoles := values.Get("includeSystemRoles") != "yes"
includeRoleObject := values.Get("includeObject") == "yes"
roles, err := clt.ListRoles(context.TODO(), &proto.ListRolesRequest{
Limit: limit,
StartKey: values.Get("startKey"),
Filter: &types.RoleFilter{
SearchKeywords: client.ParseSearchKeywords(values.Get("search"), ' '),
SkipSystemRoles: skipSystemRoles,
},
})
if err != nil {
return nil, trace.Wrap(err)
}
var typeRoles []types.Role
for _, role := range roles.GetRoles() {
typeRoles = append(typeRoles, role)
}
uiRoles, err := ui.NewRoles(typeRoles, includeRoleObject)
if err != nil {
return nil, trace.Wrap(err)
}
return &listResourcesWithoutCountGetResponse{
Items: uiRoles,
StartKey: roles.GetNextKey(),
}, nil
}
// listRequestableRolesHandle is the web handler for listing requestable roles.
// Under the hood this just calls the `ListRoles` method with a filter for requestable roles,
// we have this as a separate endpoint because the response needs to be formatted differently.
func (h *Handler) listRequestableRolesHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
rolesReq := proto.ListRequestableRolesRequest_builder{
PageSize: limit,
PageToken: values.Get("startKey"),
Filter: proto.ListRequestableRolesRequest_Filter_builder{
SearchKeywords: client.ParseSearchKeywords(values.Get("search"), ' '),
}.Build(),
}.Build()
resp, err := clt.ListRequestableRoles(r.Context(), rolesReq)
if err != nil {
return nil, trace.Wrap(err)
}
return &listResourcesWithoutCountGetResponse{
Items: ui.RequestableRolesFromProto(resp.GetRoles()),
StartKey: resp.GetNextPageToken(),
}, nil
}
func (h *Handler) deleteRole(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
roleName := params.ByName("name")
if err := clt.DeleteRole(r.Context(), roleName); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
func (h *Handler) createRoleHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
item, err := CreateResource(r, types.KindRole, services.UnmarshalRole, clt.CreateRole)
return item, trace.Wrap(err)
}
func (h *Handler) getRole(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
roleName := params.ByName("name")
role, err := clt.GetRole(r.Context(), roleName)
if err != nil {
return nil, trace.Wrap(err)
}
ri, err := ui.NewResourceItem(role)
return ri, trace.Wrap(err)
}
func (h *Handler) updateRoleHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
item, err := UpdateResource(r, params, types.KindRole, services.UnmarshalRole, clt.UpdateRole)
return item, trace.Wrap(err)
}
// getPresetRoles returns a list of preset roles expected to be available on
// this server. These are hard-coded for a given Teleport version, so this
// should have the same security implications as the Teleport version exposed
// via the public ping endpoint.
func (h *Handler) getPresetRoles(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
presets := auth.GetPresetRoles(modules.GetModules().BuildType())
return ui.NewRoles(presets, false /* without role object */)
}
// getGithubConnectorHandle returns a GitHub connector by name.
func (h *Handler) getGithubConnectorHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
connector, err := clt.GetGithubConnector(r.Context(), params.ByName("name"), true)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.NewResourceItem(connector)
}
func (h *Handler) getGithubConnectorsHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
connectors, err := getGithubConnectors(r.Context(), clt)
if err != nil {
return nil, trace.Wrap(err)
}
defaultConnectorName, defaultConnectorType, err := ProcessDefaultConnector(r.Context(), clt, connectors)
if err != nil {
return nil, trace.Wrap(err)
}
return &ui.ListAuthConnectorsResponse{
DefaultConnectorName: defaultConnectorName,
DefaultConnectorType: defaultConnectorType,
Connectors: connectors,
}, nil
}
// ProcessDefaultConnector returns the default connector type and validates that the provided connectors list contains the default connector that is set in the auth preference.
// If it isn't, it will return a fallback connector which should be used as the default, as well as update the actual auth preference to reflect the change.
func ProcessDefaultConnector(ctx context.Context, clt authclient.ClientI, connectors []ui.ResourceItem) (connectorName string, connectorType string, err error) {
authPref, err := clt.GetAuthPreference(ctx)
if err != nil {
return "", "", trace.Wrap(err, "failed to get auth preference")
}
defaultConnectorName := authPref.GetConnectorName()
defaultConnectorType := authPref.GetType()
if len(connectors) == 0 || defaultConnectorType == constants.Local {
// If there are no connectors or the default is already local, default to 'local' as the default connector.
defaultConnectorType = constants.Local
defaultConnectorName = ""
} else {
// Ensure that the default connector set in the auth preference exists in the list.
found := false
for _, c := range connectors {
if c.Name == defaultConnectorName && c.Kind == defaultConnectorType {
found = true
break
}
}
// If the default connector set in the auth preference doesn't exist, use the last connector in the list as the default.
if !found {
defaultConnectorName = connectors[len(connectors)-1].Name
defaultConnectorType = connectors[len(connectors)-1].Kind
}
}
// If the default connector we are returning here is different from the initial, also update the actual auth preference so that it's in sync.
if defaultConnectorName != authPref.GetConnectorName() || defaultConnectorType != authPref.GetType() {
authPref.SetConnectorName(defaultConnectorName)
authPref.SetType(defaultConnectorType)
_, err = clt.UpsertAuthPreference(ctx, authPref)
if err != nil {
return "", "", trace.Wrap(err, "failed to set fallback auth preference")
}
}
return defaultConnectorName, defaultConnectorType, nil
}
func getGithubConnectors(ctx context.Context, clt resourcesAPIGetter) ([]ui.ResourceItem, error) {
// TODO(okraport): DELETE IN v21.0.0, replace with regular collect.
connectors, err := clientutils.CollectWithFallback(ctx,
func(ctx context.Context, limit int, start string) ([]types.GithubConnector, string, error) {
return clt.ListGithubConnectors(ctx, limit, start, false)
},
func(ctx context.Context) ([]types.GithubConnector, error) {
return clt.GetGithubConnectors(ctx, false)
},
)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.NewGithubConnectors(connectors)
}
func (h *Handler) deleteGithubConnector(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
connectorName := params.ByName("name")
if err := clt.DeleteGithubConnector(r.Context(), connectorName); err != nil {
return nil, trace.Wrap(err)
}
authPref, err := clt.GetAuthPreference(r.Context())
if err != nil {
return nil, trace.Wrap(err, "failed to get auth preference")
}
defaultConnectorName := authPref.GetConnectorName()
defaultConnectorType := authPref.GetType()
// If the connector being deleted is the default, have the auth preference fallback to another connector.
if defaultConnectorType == constants.Github && defaultConnectorName == connectorName {
connectors, err := getGithubConnectors(r.Context(), clt)
if err != nil {
return nil, trace.Wrap(err)
}
_, _, err = ProcessDefaultConnector(r.Context(), clt, connectors)
if err != nil {
return nil, trace.Wrap(err)
}
}
return OK(), nil
}
func (h *Handler) updateGithubConnectorHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
item, err := UpdateResource[types.GithubConnector](r, params, types.KindGithubConnector, services.UnmarshalGithubConnector, clt.UpdateGithubConnector)
return item, trace.Wrap(err)
}
func (h *Handler) createGithubConnectorHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
item, err := CreateResource(r, types.KindGithubConnector, services.UnmarshalGithubConnector, clt.CreateGithubConnector)
return item, trace.Wrap(err)
}
func (h *Handler) getTrustedClustersHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
return getTrustedClusters(r.Context(), clt)
}
func getTrustedClusters(ctx context.Context, clt resourcesAPIGetter) ([]ui.ResourceItem, error) {
trustedClusters, err := stream.Collect(clientutils.Resources(ctx, clt.ListTrustedClusters))
if err != nil {
// TODO(okraport) DELETE IN v21.0.0
if trace.IsNotImplemented(err) {
trustedClusters, err = clt.GetTrustedClusters(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
} else {
return nil, trace.Wrap(err)
}
}
return ui.NewTrustedClusters(trustedClusters)
}
func (h *Handler) deleteTrustedCluster(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
tcName := params.ByName("name")
if err := clt.DeleteTrustedCluster(r.Context(), tcName); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
func (h *Handler) upsertTrustedClusterHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
var req ui.ResourceItem
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
return upsertTrustedCluster(r.Context(), clt, req.Content, r.Method, params)
}
func upsertTrustedCluster(ctx context.Context, clt resourcesAPIGetter, content, httpMethod string, params httprouter.Params) (*ui.ResourceItem, error) {
get := func(ctx context.Context, name string) (types.Resource, error) {
// Remove the MFA resp from the context before getting the trusted cluster.
// Otherwise, it will be consumed before the Upsert which actually
// requires the MFA.
// TODO(Joerger): Explicitly provide MFA response only where it is
// needed instead of removing it like this.
getCtx := mfa.ContextWithMFAResponse(ctx, nil)
return clt.GetTrustedCluster(getCtx, name)
}
extractedRes, err := ExtractResourceAndValidate(content)
if err != nil {
return nil, trace.Wrap(err)
}
if extractedRes.Kind != types.KindTrustedCluster {
return nil, trace.BadParameter("resource kind %q is invalid", extractedRes.Kind)
}
if err := CheckResourceUpsert(ctx, httpMethod, params, extractedRes.Metadata.Name, get); err != nil {
return nil, trace.Wrap(err)
}
tc, err := services.UnmarshalTrustedCluster(extractedRes.Raw)
if err != nil {
return nil, trace.Wrap(err)
}
_, err = clt.UpsertTrustedCluster(ctx, tc)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.NewResourceItem(tc)
}
// unmarshalFunc is a type signature for an unmarshaling function.
type unmarshalFunc[T types.Resource] func([]byte, ...services.MarshalOption) (T, error)
// CreateResource is a helper function for POST requests from the UI to create a new resource. It will
// validate the request contains the appropriate items, that the resource attempting to be created is
// valid. If all validations are satisfied then the creation is attempted. If the resource already exists
// a [trace.AlreadyExists] error is returned.
func CreateResource[T types.Resource](r *http.Request, kind string, unmarshalFn unmarshalFunc[T], createFn func(ctx context.Context, r T) (T, error)) (*ui.ResourceItem, error) {
var req ui.ResourceItem
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
extractedRes, err := ExtractResourceAndValidate(req.Content)
if err != nil {
return nil, trace.Wrap(err)
}
if extractedRes.Kind != kind {
return nil, trace.BadParameter("resource kind %q is invalid", extractedRes.Kind)
}
resource, err := unmarshalFn(extractedRes.Raw, services.DisallowUnknown())
if err != nil {
return nil, trace.Wrap(err)
}
created, err := createFn(r.Context(), resource)
if err != nil {
if trace.IsAlreadyExists(err) {
return nil, trace.AlreadyExists("resource with name %q already exists", extractedRes.Metadata.Name)
}
return nil, trace.Wrap(err)
}
item, err := ui.NewResourceItem(created)
return item, trace.Wrap(err)
}
// UpdateResource is a helper function for PUT requests from the UI to update an existing resource. It will
// validate the request contains the appropriate items, that the resource attempting to be updated is
// valid. If all validations are satisfied then the update is attempted. If the resource does not exist
// a [trace.NotFound] error is returned.
func UpdateResource[T types.Resource](r *http.Request, params httprouter.Params, kind string, unmarshalFn unmarshalFunc[T], updateFn func(ctx context.Context, r T) (T, error)) (*ui.ResourceItem, error) {
var req ui.ResourceItem
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
extractedRes, err := ExtractResourceAndValidate(req.Content)
if err != nil {
return nil, trace.Wrap(err)
}
if extractedRes.Kind != kind {
return nil, trace.BadParameter("resource kind %q is invalid", extractedRes.Kind)
}
resourceName := params.ByName("name")
if resourceName == "" {
return nil, trace.BadParameter("missing resource name")
}
// Error if the user is trying to rename the resource.
if extractedRes.Metadata.Name != resourceName {
return nil, trace.BadParameter("resource renaming is not supported, please create a different resource and then delete this one")
}
resource, err := unmarshalFn(extractedRes.Raw, services.DisallowUnknown())
if err != nil {
return nil, trace.Wrap(err)
}
updated, err := updateFn(r.Context(), resource)
if err != nil {
if trace.IsNotFound(err) {
return nil, trace.NotFound("resource with name %q does not exist", extractedRes.Metadata.Name)
}
return nil, trace.Wrap(err)
}
item, err := ui.NewResourceItem(updated)
return item, trace.Wrap(err)
}
// getResource tries to retrieve a resource (by name),
// returning a NotFound error if the resource does not exist.
type getResource func(context.Context, string) (types.Resource, error)
// CheckResourceUpsert checks if the resource can be created or updated, depending on the http method.
func CheckResourceUpsert(ctx context.Context, httpMethod string, params httprouter.Params, payloadResourceName string, get getResource) error {
switch httpMethod {
case http.MethodPost:
return trace.Wrap(checkResourceCreate(ctx, payloadResourceName, get))
case http.MethodPut:
resourceName := params.ByName("name")
if resourceName == "" {
return trace.BadParameter("missing resource name")
}
return trace.Wrap(checkResourceUpdate(ctx, payloadResourceName, resourceName, get))
default:
return trace.NotImplemented("http method %q not expected. this is a bug!", httpMethod)
}
}
// checkResourceCreate checks if the resource can be created, returning nil if it can.
func checkResourceCreate(ctx context.Context, payloadResourceName string, get getResource) error {
// Try to retrieve the resource by name.
_, err := get(ctx, payloadResourceName)
// If no error, then the resource already exists and cannot be created.
if err == nil {
return trace.AlreadyExists("resource with name %q already exists", payloadResourceName)
}
// If the error is not found, then the resource does not exist and can be created.
if trace.IsNotFound(err) {
return nil
}
return trace.Wrap(err)
}
// checkResourceUpdate checks if the resource can be updated, returning nil if it can.
func checkResourceUpdate(ctx context.Context, payloadResourceName, resourceName string, get getResource) error {
// Error if the user is trying to rename the resource.
if payloadResourceName != resourceName {
return trace.BadParameter("resource renaming is not supported, please create a different resource and then delete this one")
}
// Try to retrieve the resource by name.
_, err := get(ctx, payloadResourceName)
// If no error, then the resource already exists and can be updated.
if err == nil {
return nil
}
// If the error is not found, then the resource does not exist and cannot be updated.
if trace.IsNotFound(err) {
return trace.NotFound("resource with name %q does not exist", payloadResourceName)
}
return trace.Wrap(err)
}
// ExtractResourceAndValidate extracts resource information from given string and validates basic fields.
func ExtractResourceAndValidate(yaml string) (*services.UnknownResource, error) {
unknownRes, err := extractResource(yaml)
if err != nil {
return nil, trace.Wrap(err)
}
if err := unknownRes.Metadata.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &unknownRes, nil
}
func extractResource(yaml string) (services.UnknownResource, error) {
var unknownRes services.UnknownResource
reader := strings.NewReader(yaml)
decoder := kyaml.NewYAMLOrJSONDecoder(reader, 32*1024)
if err := decoder.Decode(&unknownRes); err != nil {
return services.UnknownResource{}, trace.BadParameter("not a valid resource declaration")
}
return unknownRes, nil
}
func convertListResourcesRequest(r *http.Request, kind string) (*proto.ListResourcesRequest, error) {
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
sortBy := types.GetSortByFromString(values.Get("sort"))
startKey := values.Get("startKey")
return &proto.ListResourcesRequest{
ResourceType: kind,
Limit: limit,
StartKey: startKey,
SortBy: sortBy,
PredicateExpression: values.Get("query"),
SearchKeywords: client.ParseSearchKeywords(values.Get("search"), ' '),
UseSearchAsRoles: values.Get("searchAsRoles") == "yes",
}, nil
}
// listKubeResources gets a list of kubernetes resources depending on the type of resource.
func listKubeResources(ctx context.Context, kubeClient kubeproto.KubeServiceClient, values url.Values, site, resourceKind string) (*kubeproto.ListKubernetesResourcesResponse, error) {
req, err := newKubeListRequest(values, site, resourceKind)
if err != nil {
return nil, trace.Wrap(err)
}
return kubeClient.ListKubernetesResources(ctx, req)
}
// newKubeListRequest parses the request parameters into a ListKubernetesResourcesRequest.
func newKubeListRequest(values url.Values, site, resourceKind string) (*kubeproto.ListKubernetesResourcesRequest, error) {
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
sortBy := types.GetSortByFromString(values.Get("sort"))
startKey := values.Get("startKey")
req := kubeproto.ListKubernetesResourcesRequest_builder{
ResourceType: resourceKind,
Limit: limit,
StartKey: startKey,
SortBy: &sortBy,
PredicateExpression: values.Get("query"),
SearchKeywords: client.ParseSearchKeywords(values.Get("search"), ' '),
UseSearchAsRoles: values.Get("searchAsRoles") == "yes",
TeleportCluster: site,
KubernetesCluster: values.Get("kubeCluster"),
KubernetesNamespace: values.Get("kubeNamespace"),
}.Build()
return req, nil
}
type listResourcesGetResponse struct {
// Items is a list of resources retrieved.
Items any `json:"items"`
// StartKey is the position to resume search events.
StartKey string `json:"startKey"`
// TotalCount is the total count of resources available
// after filter.
TotalCount int `json:"totalCount"`
}
type listResourcesWithoutCountGetResponse struct {
// Items is a list of resources retrieved.
Items any `json:"items"`
// StartKey is the position to resume search events.
StartKey string `json:"startKey"`
}
type resourcesAPIGetter interface {
// GetRole returns role by name
GetRole(ctx context.Context, name string) (types.Role, error)
// ListRoles returns a paginated list of roles.
ListRoles(ctx context.Context, req *proto.ListRolesRequest) (*proto.ListRolesResponse, error)
// ListRequestableRoles returns a paginated list of requestable roles.
ListRequestableRoles(ctx context.Context, req *proto.ListRequestableRolesRequest) (*proto.ListRequestableRolesResponse, error)
// UpsertRole creates or updates role
UpsertRole(ctx context.Context, role types.Role) (types.Role, error)
// GetGithubConnectors returns all configured Github connectors
GetGithubConnectors(ctx context.Context, withSecrets bool) ([]types.GithubConnector, error)
// ListGithubConnectors returns a page of valid registered Github connectors.
// withSecrets adds or removes client secret from return results.
ListGithubConnectors(ctx context.Context, limit int, start string, withSecrets bool) ([]types.GithubConnector, string, error)
// RangeGithubConnectors returns valid registered Github connectors within the range [start, end).
// withSecrets adds or removes client secret from return results.
RangeGithubConnectors(ctx context.Context, start, end string, withSecrets bool) iter.Seq2[types.GithubConnector, error]
// GetGithubConnector returns the specified Github connector
GetGithubConnector(ctx context.Context, id string, withSecrets bool) (types.GithubConnector, error)
// DeleteGithubConnector deletes the specified Github connector
DeleteGithubConnector(ctx context.Context, id string) error
// UpsertTrustedCluster creates or updates a TrustedCluster in the backend.
UpsertTrustedCluster(ctx context.Context, tc types.TrustedCluster) (types.TrustedCluster, error)
// GetTrustedCluster returns a single TrustedCluster by name.
GetTrustedCluster(ctx context.Context, name string) (types.TrustedCluster, error)
// GetTrustedClusters returns all TrustedClusters in the backend.
GetTrustedClusters(ctx context.Context) ([]types.TrustedCluster, error)
// ListTrustedClusters returns a page of Trusted Cluster resources.
ListTrustedClusters(ctx context.Context, limit int, startKey string) ([]types.TrustedCluster, string, error)
// DeleteTrustedCluster removes a TrustedCluster from the backend by name.
DeleteTrustedCluster(ctx context.Context, name string) error
// ListResources returns a paginated list of resources.
ListResources(ctx context.Context, req proto.ListResourcesRequest) (*types.ListResourcesResponse, error)
}
/*
* Teleport
* Copyright (C) 2025 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"
"fmt"
"net/http"
"os"
"strconv"
"github.com/coreos/go-semver/semver"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/utils/teleportassets"
"github.com/gravitational/teleport/lib/web/scripts"
)
const (
insecureParamName = "insecure"
groupParamName = "group"
)
// installScriptHandle handles calls for "/scripts/install.sh" and responds with a bash script installing Teleport
// by downloading and running `teleport-update`. This installation script does not start the agent, join it,
// or configure its services. This is handled by the "/scripts/:token/install-*.sh" scripts.
func (h *Handler) installScriptHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params) (any, error) {
// This is a hack because the router is not allowing us to register "/scripts/install.sh", so we use
// the parameter ":token" to match the script name.
// Currently, only "install.sh" is supported.
if params.ByName("token") != "install.sh" {
return nil, trace.NotFound(`Route not found, query "/scripts/install.sh" for the install-only script, or "/scripts/:token/install-node.sh" for the install + join script.`)
}
// TODO(hugoShaka): cache function
opts, err := h.installScriptOptions(r.Context())
if err != nil {
return nil, trace.Wrap(err, "Failed to build install script options")
}
if insecure := r.URL.Query().Get(insecureParamName); insecure != "" {
v, err := strconv.ParseBool(insecure)
if err != nil {
return nil, trace.BadParameter("failed to parse insecure flag %q: %v", insecure, err)
}
opts.Insecure = v
}
if group := r.URL.Query().Get(groupParamName); group != "" {
opts.Group = group
}
script, err := scripts.GetInstallScript(r.Context(), opts)
if err != nil {
h.logger.WarnContext(r.Context(), "Failed to get install script", "error", err)
return nil, trace.Wrap(err, "getting script")
}
w.WriteHeader(http.StatusOK)
if _, err := fmt.Fprintln(w, script); err != nil {
h.logger.WarnContext(r.Context(), "Failed to write install script", "error", err)
}
return nil, nil
}
// installScriptOptions computes the agent installation options based on the proxy configuration and the cluster status.
// This includes:
// - the type of automatic updates
// - the desired version
// - the proxy address (used for updates).
// - the Teleport artifact name and CDN
func (h *Handler) installScriptOptions(ctx context.Context) (scripts.InstallScriptOptions, error) {
const defaultGroup, defaultUpdater = "", ""
version, err := h.autoUpdateResolver.GetVersion(ctx, defaultGroup, defaultUpdater)
if err != nil {
h.logger.WarnContext(ctx, "Failed to get intended agent version", "error", err)
version = teleport.SemVer()
}
// if there's a rollout, we do new autoupdates
_, rolloutErr := h.cfg.AccessPoint.GetAutoUpdateAgentRollout(ctx)
if rolloutErr != nil && !trace.IsNotFound(rolloutErr) && !trace.IsNotImplemented(rolloutErr) {
h.logger.WarnContext(ctx, "Failed to get rollout", "error", rolloutErr)
return scripts.InstallScriptOptions{}, trace.Wrap(err, "failed to check the autoupdate agent rollout state")
}
var autoupdateStyle scripts.AutoupdateStyle
switch {
case rolloutErr == nil:
autoupdateStyle = scripts.UpdaterBinaryAutoupdate
case automaticUpgrades(h.GetClusterFeatures()):
autoupdateStyle = scripts.PackageManagerAutoupdate
default:
autoupdateStyle = scripts.NoAutoupdate
}
var teleportFlavor string
switch h.cfg.Modules.BuildType() {
case modules.BuildEnterprise:
teleportFlavor = types.PackageNameEnt
case modules.BuildOSS, modules.BuildCommunity:
teleportFlavor = types.PackageNameOSS
default:
h.logger.WarnContext(ctx, "Unknown built type, defaulting to the 'teleport' package.", "type", h.cfg.Modules.BuildType())
teleportFlavor = types.PackageNameOSS
}
cdnBaseURL, err := getCDNBaseURL(h.cfg.Modules.BuildType(), version)
if err != nil {
h.logger.WarnContext(ctx, "Failed to get CDN base URL", "error", err)
return scripts.InstallScriptOptions{}, trace.Wrap(err)
}
return scripts.InstallScriptOptions{
AutoupdateStyle: autoupdateStyle,
TeleportVersion: version,
CDNBaseURL: cdnBaseURL,
ProxyAddr: h.PublicProxyAddr(),
TeleportFlavor: teleportFlavor,
FIPS: modules.IsFIPSBuild(),
}, nil
}
// EnvVarCDNBaseURL is the environment variable that allows users to override the Teleport base CDN url used in the installation script.
// Setting this value is required for testing (make production builds install from the dev CDN, and vice versa).
// As we (the Teleport company) don't distribute AGPL binaries, this must be set when using a Teleport OSS build.
// Example values:
// - "https://cdn.teleport.dev" (prod)
// - "https://cdn.cloud.gravitational.io" (dev builds/staging)
const EnvVarCDNBaseURL = "TELEPORT_CDN_BASE_URL"
func getCDNBaseURL(buildType string, version *semver.Version) (string, error) {
// If the user explicitly overrides the CDN base URL, we use it.
if override := os.Getenv(EnvVarCDNBaseURL); override != "" {
return override, nil
}
// If this is an AGPL build, we don't want to automatically install binaries distributed under a more restrictive
// license so we error and ask the user set the CDN URL, either to:
// - the official Teleport CDN if they agree with the community license and meet its requirements
// - a custom CDN where they can store their own AGPL binaries
if buildType == modules.BuildOSS {
return "", trace.BadParameter(
"This proxy is licensed under AGPL but CDN binaries are licensed under the more restrictive Community license. "+
"You can set TELEPORT_CDN_BASE_URL to a custom CDN, or to %q if you are OK with using the Community Edition license.",
teleportassets.CDNBaseURL())
}
return teleportassets.CDNBaseURLForVersion(version), nil
}
/*
* 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"
"log/slog"
"net"
"net/http"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/lib/defaults"
)
// ServerConfig provides dependencies required to create a [Server].
type ServerConfig struct {
// Server serves the web api
Server *http.Server
// Handler web handler
Handler *APIHandler
// Log to write log messages
Log *slog.Logger
// ShutdownPollPeriod sets polling period for shutdown
ShutdownPollPeriod time.Duration
}
// CheckAndSetDefaults validates fields and populates empty fields with default values.
func (c *ServerConfig) CheckAndSetDefaults() error {
if c.Server == nil {
return trace.BadParameter("missing required parameter Server")
}
if c.Handler == nil {
return trace.BadParameter("missing required parameter Handler")
}
if c.ShutdownPollPeriod <= 0 {
c.ShutdownPollPeriod = defaults.ShutdownPollPeriod
}
if c.Log == nil {
c.Log = slog.With(teleport.ComponentKey, teleport.ComponentProxy)
}
return nil
}
// Server serves the web api.
type Server struct {
cfg ServerConfig
mu sync.Mutex
ln net.Listener
closed bool
}
// NewServer constructs a [Server] from the provided [ServerConfig].
func NewServer(cfg ServerConfig) (*Server, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &Server{
cfg: cfg,
}, nil
}
// Serve launches the configured [http.Server].
func (s *Server) Serve(l net.Listener) error {
s.mu.Lock()
s.ln = l
closed := s.closed
if closed {
s.ln.Close()
}
s.mu.Unlock()
if closed {
return trace.Errorf("serve called on previously closed server")
}
return trace.Wrap(s.cfg.Server.Serve(l))
}
// Close immediately closes the [http.Server].
func (s *Server) Close() error {
s.mu.Lock()
s.closed = true
if s.ln != nil {
s.ln.Close()
}
s.mu.Unlock()
return trace.NewAggregate(s.cfg.Handler.Close(), s.cfg.Server.Close())
}
// HandleConnection handles connections from plain TCP applications.
func (s *Server) HandleConnection(ctx context.Context, conn net.Conn) error {
return s.cfg.Handler.appHandler.HandleConnection(ctx, conn)
}
// Shutdown initiates graceful shutdown. The underlying [http.Server]
// is not shutdown until all active connections are terminated or
// the context times out. This is required because the [http.Server]
// does not attempt to close nor wait for hijacked connections such as
// WebSockets during Shutdown; which means that any open sessions in the
// web UI will not prevent the [http.Server] from shutting down.
func (s *Server) Shutdown(ctx context.Context) error {
s.mu.Lock()
var err error
s.closed = true
if s.ln != nil {
err = s.ln.Close()
}
s.mu.Unlock()
activeConnections := s.cfg.Handler.handler.userConns.Load()
if activeConnections == 0 {
err := s.cfg.Server.Shutdown(ctx)
return trace.NewAggregate(err, s.cfg.Handler.Close())
}
s.cfg.Log.InfoContext(ctx, "Shutdown: waiting for active connections to finish", "active_connection_count", activeConnections)
lastReport := time.Time{}
ticker := time.NewTicker(s.cfg.ShutdownPollPeriod)
defer ticker.Stop()
for {
select {
case <-ticker.C:
activeConnections = s.cfg.Handler.handler.userConns.Load()
if activeConnections == 0 {
err := s.cfg.Server.Shutdown(ctx)
return trace.NewAggregate(err, s.cfg.Handler.Close())
}
if time.Since(lastReport) > 10*s.cfg.ShutdownPollPeriod {
s.cfg.Log.InfoContext(ctx, "Shutdown: waiting for active connections to finish", "active_connection_count", activeConnections)
lastReport = time.Now()
}
case <-ctx.Done():
s.cfg.Log.InfoContext(ctx, "Context canceled wait, returning")
return trace.ConnectionProblem(trace.NewAggregate(err, s.cfg.Handler.Close()), "context canceled")
}
}
}
/*
* 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 (
"net/http"
"slices"
"strings"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client"
linuxdesktopv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/linuxdesktop/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/ui"
"github.com/gravitational/teleport/lib/utils/set"
webui "github.com/gravitational/teleport/lib/web/ui"
)
// clusterKubesGet returns a list of kube clusters in a form the UI can present.
func (h *Handler) clusterKubesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindKubernetesCluster)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetResourcePage[types.KubeCluster](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: webui.MakeKubeClusters(page.Resources, accessChecker),
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
// clusterKubeResourcesGet returns supported requested kubernetes subresources eg: pods, namespaces, secrets etc.
func (h *Handler) clusterKubeResourcesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
kind := r.URL.Query().Get("kind")
kubeCluster := r.URL.Query().Get("kubeCluster")
if kubeCluster == "" {
return nil, trace.BadParameter("missing param %q", "kubeCluster")
}
if kind == "" {
return nil, trace.BadParameter("missing param %q", "kind")
}
if !slices.Contains(types.KubernetesResourcesKinds, kind) && !strings.HasPrefix(kind, types.AccessRequestPrefixKindKube) {
return nil, trace.BadParameter("kind is not valid, valid kinds %v %s<kind>", types.KubernetesResourcesKinds, types.AccessRequestPrefixKindKube)
}
clt, err := sctx.NewKubernetesServiceClient(r.Context(), h.cfg.ProxyWebAddr.Addr)
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := listKubeResources(r.Context(), clt, r.URL.Query(), cluster.GetName(), kind)
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: webui.MakeKubeResources(resp.GetResources(), kubeCluster),
StartKey: resp.GetNextKey(),
TotalCount: int(resp.GetTotalCount()),
}, nil
}
// clusterKubeServersList returns a list of kube servers in a form the UI can present.
func (h *Handler) clusterKubeServersList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := ctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindKubeServer)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetResourcePage[types.KubeServer](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: page.Resources,
StartKey: page.NextKey,
}, nil
}
// clusterDatabasesGet returns a list of db servers in a form the UI can present.
func (h *Handler) clusterDatabasesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindDatabaseServer)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetResourcePage[types.DatabaseServer](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
uiItems := make([]webui.Database, 0, len(page.Resources))
for _, dbServer := range page.Resources {
db := webui.MakeDatabaseFromDatabaseServer(dbServer, accessChecker, h.cfg.DatabaseREPLRegistry, false /* requires reset*/)
uiItems = append(uiItems, db)
}
return listResourcesGetResponse{
Items: uiItems,
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
// clusterDatabaseGet returns a database in a form the UI can present.
func (h *Handler) clusterDatabaseGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
databaseName := p.ByName("database")
if databaseName == "" {
return nil, trace.BadParameter("database name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
dbServers, err := fetchDatabaseServersWithName(r.Context(), clt, r, databaseName)
if err != nil {
return nil, trace.Wrap(err)
}
aggregateStatus := types.AggregateHealthStatus(func(yield func(types.TargetHealthStatus) bool) {
for _, srv := range dbServers {
if !yield(srv.GetTargetHealthStatus()) {
return
}
}
})
dbServers[0].SetTargetHealthStatus(aggregateStatus)
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
return webui.MakeDatabaseFromDatabaseServer(
dbServers[0],
accessChecker,
h.cfg.DatabaseREPLRegistry,
false, /* requiresRequest */
), nil
}
// clusterDatabaseServicesList returns a list of DatabaseServices (database agents) in a form the UI can present.
func (h *Handler) clusterDatabaseServicesList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := ctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindDatabaseService)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetResourcePage[types.DatabaseService](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: webui.MakeDatabaseServices(page.Resources),
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
// clusterDatabaseServersList returns a list of database servers in a form the UI can present.
func (h *Handler) clusterDatabaseServersList(w http.ResponseWriter, r *http.Request, p httprouter.Params, ctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := ctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindDatabaseServer)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetResourcePage[types.DatabaseServer](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: page.Resources,
StartKey: page.NextKey,
}, nil
}
// clusterDesktopsGet returns a list of desktops in a form the UI can present.
func (h *Handler) clusterDesktopsGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindWindowsDesktop)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetEnrichedResourcePage(r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
uiDesktops := make([]webui.Desktop, 0, len(page.Resources))
for _, r := range page.Resources {
switch desktop := r.ResourceWithLabels.(type) {
case types.WindowsDesktop:
uiDesktops = append(uiDesktops, webui.MakeWindowsDesktop(desktop, r.Logins, false /* requiresRequest */))
case types.Resource153UnwrapperT[*linuxdesktopv1.LinuxDesktop]:
uiDesktops = append(uiDesktops, webui.MakeLinuxDesktop(desktop.UnwrapT(), r.Logins, false /* requiresRequest */))
}
}
return listResourcesGetResponse{
Items: uiDesktops,
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
// clusterDesktopServicesGet returns a list of desktop services in a form the UI can present.
func (h *Handler) clusterDesktopServicesGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
// Get a client to the Auth Server with the logged in user's identity. The
// identity of the logged in user is used to fetch the list of desktop services.
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindWindowsDesktopService)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := client.GetResourcePage[types.WindowsDesktopService](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: webui.MakeDesktopServices(page.Resources),
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
// getDesktopHandle returns a desktop.
func (h *Handler) getDesktopHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
desktopName := p.ByName("desktopName")
windowsDesktops, err := clt.GetWindowsDesktops(r.Context(), types.WindowsDesktopFilter{Name: desktopName})
if err != nil {
return nil, trace.Wrap(err)
}
if len(windowsDesktops) == 0 {
return nil, trace.NotFound("expected at least 1 desktop, got 0")
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
// windowsDesktops may contain the same desktop multiple times
// if multiple Windows Desktop Services are in use. We only need
// to see the desktop once in the UI, so just take the first one.
desktop := windowsDesktops[0]
logins, err := accessChecker.GetAllowedLoginsForResource(desktop)
if err != nil {
return nil, trace.Wrap(err)
}
return webui.MakeWindowsDesktop(desktop, logins, false /* requiresRequest */), nil
}
// desktopIsActive checks if a desktop has an active session and returns a desktopIsActive.
//
// GET /v1/webapi/sites/:site/desktops/:desktopName/active
//
// Response body:
//
// {"active": bool}
func (h *Handler) desktopIsActive(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
desktopName := p.ByName("desktopName")
trackers, err := h.auth.proxyClient.GetActiveSessionTrackersWithFilter(r.Context(), &types.SessionTrackerFilter{
Kind: string(types.WindowsDesktopSessionKind),
State: &types.NullableSessionState{
State: types.SessionState_SessionStateRunning,
},
DesktopName: desktopName,
})
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
for _, tracker := range trackers {
// clt is an auth.ClientI with the role of the user, so
// clt.GetWindowsDesktops() can be used to confirm that
// the user has access to the requested desktop.
desktops, err := clt.GetWindowsDesktops(r.Context(),
types.WindowsDesktopFilter{Name: tracker.GetDesktopName()})
if err != nil {
return nil, trace.Wrap(err)
}
if len(desktops) == 0 {
// There are no active sessions for this desktop
// or the user doesn't have access to it
break
} else {
return desktopIsActive{true}, nil
}
}
return desktopIsActive{false}, nil
}
type desktopIsActive struct {
Active bool `json:"active"`
}
// createNodeRequest contains the required information to create a Node.
type createNodeRequest struct {
Name string `json:"name,omitempty"`
SubKind string `json:"subKind,omitempty"`
Hostname string `json:"hostname,omitempty"`
Addr string `json:"addr,omitempty"`
Labels []ui.Label `json:"labels,omitempty"`
AWSInfo *webui.AWSMetadata `json:"aws,omitempty"`
}
func (r *createNodeRequest) checkAndSetDefaults() error {
if r.Name == "" {
return trace.BadParameter("missing node name")
}
// Nodes provided by the Teleport Agent are not meant to be created by the user.
// They connect to the cluster and heartbeat their information.
//
// Agentless Nodes with Teleport CA call the Teleport Proxy and upsert themselves,
// so they are also not meant to be added from web api.
if r.SubKind != types.SubKindOpenSSHEICENode {
return trace.BadParameter("invalid subkind %q, only %q is supported", r.SubKind, types.SubKindOpenSSHEICENode)
}
if r.Hostname == "" {
return trace.BadParameter("missing node hostname")
}
if r.Addr == "" {
return trace.BadParameter("missing node addr")
}
return nil
}
// handleNodeCreate creates a Teleport Node.
func (h *Handler) handleNodeCreate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
ctx := r.Context()
var req *createNodeRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
labels := make(map[string]string, len(req.Labels))
for _, label := range req.Labels {
labels[label.Name] = label.Value
}
server, err := types.NewNode(
req.Name,
req.SubKind,
types.ServerSpecV2{
Hostname: req.Hostname,
Addr: req.Addr,
CloudMetadata: &types.CloudMetadata{
AWS: &types.AWSInfo{
AccountID: req.AWSInfo.AccountID,
InstanceID: req.AWSInfo.InstanceID,
Region: req.AWSInfo.Region,
VPCID: req.AWSInfo.VPCID,
Integration: req.AWSInfo.Integration,
SubnetID: req.AWSInfo.SubnetID,
},
},
},
labels,
)
if err != nil {
return nil, trace.Wrap(err)
}
if _, err := clt.UpsertNode(r.Context(), server); err != nil {
return nil, trace.Wrap(err)
}
accessChecker, err := sctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
logins, err := accessChecker.GetAllowedLoginsForResource(server)
if err != nil {
return nil, trace.Wrap(err)
}
loginSet := set.New(logins...)
return webui.MakeServer(server, webui.MakeServerConfig{
ClusterName: cluster.GetName(),
Logins: &webui.PrincipalSet{All: loginSet, Granted: loginSet},
RequiresRequest: false,
SupportedFeatures: nil,
}), nil
}
/*
* 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"
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"io"
"log/slog"
"net"
"slices"
"sync"
"time"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
"go.opentelemetry.io/contrib/instrumentation/google.golang.org/grpc/otelgrpc"
"golang.org/x/crypto/ssh"
"golang.org/x/crypto/ssh/agent"
"golang.org/x/net/http2"
"golang.org/x/sync/singleflight"
"google.golang.org/grpc"
"google.golang.org/grpc/credentials"
"github.com/gravitational/teleport"
"github.com/gravitational/teleport/api/breaker"
apiclient "github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
kubeproto "github.com/gravitational/teleport/api/gen/proto/go/teleport/kube/v1"
"github.com/gravitational/teleport/api/metadata"
"github.com/gravitational/teleport/api/types"
apiutils "github.com/gravitational/teleport/api/utils"
apisshutils "github.com/gravitational/teleport/api/utils/sshutils"
"github.com/gravitational/teleport/lib/auth/authclient"
"github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/authz"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/modules"
"github.com/gravitational/teleport/lib/multiplexer"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/services"
alpncommon "github.com/gravitational/teleport/lib/srv/alpnproxy/common"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/sshutils"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
// SessionContext is a context associated with a user's
// web session. An instance of the context is created for
// each web session generated for the user and provides
// a basic client cache for remote auth server connections.
type SessionContext struct {
// SessionContextConfig contains dependency injected configurations
cfg SessionContextConfig
// remoteClientCache holds the remote clients that have been used in this
// session.
remoteClientCache
// remoteClientGroup prevents duplicate requests to create remote clients
// for a given cluster
remoteClientGroup singleflight.Group
// mu guards kubeGRPCServiceConn.
mu sync.Mutex
// kubeGRPCServiceConn is a connection to the kubernetes service.
kubeGRPCServiceConn *grpc.ClientConn
}
type SessionContextConfig struct {
// Log is used to emit logs
Log *slog.Logger
// User is the name of the current user
User string
// RootClusterName is the name of the root cluster
RootClusterName string
// RootClient holds a connection to the root auth. Note that requests made using this
// client are made with the identity of the user and are NOT cached.
RootClient *authclient.Client
// UnsafeCachedAuthClient holds a read-only cache to root auth. Note this access
// point cache is authenticated with the identity of the node, not of the
// user. This is why its prefixed with "unsafe".
//
// This access point should only be used if the identity of the caller will
// not affect the result of the RPC. For example, never use it to call
// "GetNodes".
UnsafeCachedAuthClient authclient.ReadProxyAccessPoint
// UnsafeScoopedRoleReader is a scoped role reader. It's authenticated with
// the identity of the node, hence it's "unsafe".
//
// Only use it where the identity of the caller will not affect the result of
// the RPC. For example, never call it to get a list of roles and return it
// to the user, as it may return more than the user is allowed to see.
UnsafeScopedRoleReader services.ScopedRoleReader
Parent *sessionCache
// Resources is a persistent resource store this context is bound to.
// The store maintains a list of resources between session renewals
Resources *sessionResources
// Session refers the web session created for the user.
Session types.WebSession
// newRemoteClient is used by tests to override how remote clients are constructed to allow for fake clusters
newRemoteClient func(ctx context.Context, sessionContext *SessionContext, cluster reversetunnelclient.Cluster) (authclient.ClientI, error)
}
func (c *SessionContextConfig) CheckAndSetDefaults() error {
if c.RootClient == nil {
return trace.BadParameter("RootClient required")
}
if c.UnsafeCachedAuthClient == nil {
return trace.BadParameter("UnsafeCachedAuthClient required")
}
if c.Parent == nil {
return trace.BadParameter("Parent required")
}
if c.Resources == nil {
return trace.BadParameter("Resources required")
}
if c.Session == nil {
return trace.BadParameter("Session required")
}
if c.UnsafeCachedAuthClient == nil {
return trace.BadParameter("Scoped role reader required")
}
if c.Log == nil {
c.Log = slog.With(
"user", c.User,
"session", c.Session.GetShortName(),
)
}
if c.newRemoteClient == nil {
c.newRemoteClient = newRemoteClient
}
if c.RootClusterName == "" {
c.RootClusterName = c.Parent.clusterName
}
return nil
}
func NewSessionContext(cfg SessionContextConfig) (*SessionContext, error) {
if err := cfg.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
return &SessionContext{
cfg: cfg,
}, nil
}
// String returns the text representation of this context
func (c *SessionContext) String() string {
return fmt.Sprintf("WebSession(user=%v,id=%v,expires=%v,bearer_expires=%v)",
c.cfg.User,
c.cfg.Session.GetShortName(),
c.cfg.Session.GetExpiryTime(),
c.cfg.Session.GetBearerTokenExpiryTime(),
)
}
// AddClosers adds the specified closers to this context
func (c *SessionContext) AddClosers(closers ...io.Closer) {
c.cfg.Resources.addClosers(closers...)
}
// RemoveCloser removes the specified closer from this context
func (c *SessionContext) RemoveCloser(closer io.Closer) {
c.cfg.Resources.removeCloser(closer)
}
// Invalidate invalidates this context by removing the underlying session
// and closing all underlying closers
func (c *SessionContext) Invalidate(ctx context.Context) error {
return c.cfg.Parent.invalidateSession(ctx, c)
}
func (c *SessionContext) validateBearerToken(ctx context.Context, token string) error {
fetchedToken, err := c.cfg.Parent.readBearerToken(ctx, types.GetWebTokenRequest{
User: c.cfg.User,
Token: token,
})
if err != nil {
return trace.Wrap(err)
}
if fetchedToken.GetUser() != c.cfg.User {
c.cfg.Log.WarnContext(ctx, "Failed validating bearer token: the user in bearer token did not match the user for session",
"token_user", fetchedToken.GetUser(),
"token", token,
"session_user", c.cfg.User,
"session_id", c.GetSessionID(),
)
return trace.AccessDenied("access denied")
}
return nil
}
// GetClient returns the client connected to the auth server
func (c *SessionContext) GetClient() (authclient.ClientI, error) {
return c.cfg.RootClient, nil
}
// GetClientConnection returns a connection to Auth Service
func (c *SessionContext) GetClientConnection() *grpc.ClientConn {
return c.cfg.RootClient.GetConnection()
}
// GetUserClient will return an [authclient.ClientI] with the role of the user at
// the requested cluster. If the cluster is local a client with the users local role
// is returned. If the cluster is remote a client with the users remote role is
// returned.
func (c *SessionContext) GetUserClient(ctx context.Context, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
// if we're trying to access the local cluster, pass back the local client.
if c.cfg.RootClusterName == cluster.GetName() {
return c.cfg.RootClient, nil
}
// return the client for the requested remote cluster
clt, err := c.remoteClient(ctx, cluster)
return clt, trace.Wrap(err)
}
// remoteClient returns an [authclient.ClientI] with the role of the user at
// the requested [cluster]. All remote clients are lazily created
// when they are first requested and then cached. Subsequent requests
// will return the previously created client to prevent having more than
// a single [authclient.ClientI] per cluster for a user.
//
// A [singleflight.Group] is leveraged to prevent duplicate requests for remote
// clients at the same time to race.
func (c *SessionContext) remoteClient(ctx context.Context, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
cltI, err, _ := c.remoteClientGroup.Do(cluster.GetName(), func() (any, error) {
// check if we already have a connection to this cluster
if clt, ok := c.remoteClientCache.getRemoteClient(cluster); ok {
return clt, nil
}
rClt, err := c.cfg.newRemoteClient(ctx, c, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
// we'll save the remote client in our session context so we don't have to
// build a new connection next time. all remote clients will be closed when
// the session context is closed.
err = c.remoteClientCache.addRemoteClient(cluster, rClt)
if err != nil {
c.cfg.Log.InfoContext(ctx, "Failed closing stale remote client for cluster",
"remote_cluster", cluster.GetName(),
"error", err,
)
}
return rClt, nil
})
if err != nil {
return nil, trace.Wrap(err)
}
clt, ok := cltI.(authclient.ClientI)
if !ok {
return nil, trace.BadParameter("unexpected type %T received for auth client", cltI)
}
return clt, nil
}
// newRemoteClient returns a client to a remote cluster with the role of current user.
func newRemoteClient(ctx context.Context, sctx *SessionContext, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
clt, err := sctx.newRemoteTLSClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
// Clients lazily dial, so attempt an RPC to determine if this client
// is functional or not.
_, err = clt.GetDomainName(ctx)
if err != nil {
return nil, trace.NewAggregate(err, clt.Close())
}
return clt, nil
}
// clusterDialer returns DialContext function using cluster's dial function
func clusterDialer(remoteCluster reversetunnelclient.Cluster, src, dst net.Addr) apiclient.ContextDialer {
return apiclient.ContextDialerFunc(func(in context.Context, network, _ string) (net.Conn, error) {
dialParams := reversetunnelclient.DialParams{
From: src,
OriginalClientDstAddr: dst,
}
clientSrcAddr, clientDstAddr := authz.ClientAddrsFromContext(in)
if dialParams.From == nil && clientSrcAddr != nil {
dialParams.From = clientSrcAddr
}
if dialParams.OriginalClientDstAddr == nil && clientDstAddr != nil {
dialParams.OriginalClientDstAddr = clientDstAddr
}
return remoteCluster.DialAuthServer(dialParams)
})
}
// NewKubernetesServiceClient returns a new KubernetesServiceClient.
func (c *SessionContext) NewKubernetesServiceClient(ctx context.Context, addr string) (kubeproto.KubeServiceClient, error) {
c.mu.Lock()
conn := c.kubeGRPCServiceConn
c.mu.Unlock()
if conn != nil {
return kubeproto.NewKubeServiceClient(conn), nil
}
tlsConfig, err := c.ClientTLSConfig(ctx, c.cfg.RootClusterName)
if err != nil {
return nil, trace.Wrap(err)
}
// Set the ALPN protocols to use when dialing the proxy gRPC mTLS endpoint.
tlsConfig.NextProtos = []string{string(alpncommon.ProtocolProxyGRPCSecure), http2.NextProtoTLS}
conn, err = grpc.DialContext(
ctx,
addr,
grpc.WithTransportCredentials(credentials.NewTLS(tlsConfig)),
grpc.WithStatsHandler(otelgrpc.NewClientHandler()),
grpc.WithChainUnaryInterceptor(
metadata.UnaryClientInterceptor,
),
grpc.WithChainStreamInterceptor(
metadata.StreamClientInterceptor,
),
)
if err != nil {
return nil, trace.Wrap(err)
}
c.mu.Lock()
c.kubeGRPCServiceConn = conn
c.mu.Unlock()
return kubeproto.NewKubeServiceClient(conn), nil
}
// ClientTLSConfig returns client TLS authentication associated
// with the web session context
func (c *SessionContext) ClientTLSConfig(ctx context.Context, clusterName ...string) (*tls.Config, error) {
var certPool *x509.CertPool
if len(clusterName) == 0 {
certAuthorities, err := c.cfg.Parent.proxyClient.GetCertAuthorities(ctx, types.HostCA, false)
if err != nil {
return nil, trace.Wrap(err)
}
certPool, _, err = services.CertPoolFromCertAuthorities(certAuthorities)
if err != nil {
return nil, trace.Wrap(err)
}
} else {
certAuthority, err := c.cfg.Parent.proxyClient.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: clusterName[0],
}, false)
if err != nil {
return nil, trace.Wrap(err)
}
certPool, err = services.CertPool(certAuthority)
if err != nil {
return nil, trace.Wrap(err)
}
}
tlsConfig := utils.TLSConfig(c.cfg.Parent.cipherSuites)
tlsCert, err := tls.X509KeyPair(c.cfg.Session.GetTLSCert(), c.cfg.Session.GetTLSPriv())
if err != nil {
return nil, trace.Wrap(err, "failed to parse TLS cert and key")
}
tlsConfig.Certificates = []tls.Certificate{tlsCert}
tlsConfig.RootCAs = certPool
tlsConfig.ServerName = apiutils.EncodeClusterName(c.cfg.Parent.clusterName)
tlsConfig.Time = c.cfg.Parent.clock.Now
return tlsConfig, nil
}
func (c *SessionContext) newRemoteTLSClient(ctx context.Context, cluster reversetunnelclient.Cluster) (authclient.ClientI, error) {
tlsConfig, err := c.ClientTLSConfig(ctx, cluster.GetName())
if err != nil {
return nil, trace.Wrap(err)
}
clientSrcAddr, clientDstAddr := authz.ClientAddrsFromContext(ctx)
return authclient.NewClient(apiclient.Config{
Context: ctx,
Dialer: clusterDialer(cluster, clientSrcAddr, clientDstAddr),
Credentials: []apiclient.Credentials{
apiclient.LoadTLS(tlsConfig),
},
CircuitBreakerConfig: breaker.NoopBreakerConfig(),
})
}
// GetUser returns the authenticated teleport user
func (c *SessionContext) GetUser() string {
return c.cfg.User
}
// extendWebSession creates a new web session for this user
// based on the previous session
func (c *SessionContext) extendWebSession(ctx context.Context, req renewSessionRequest) (types.WebSession, error) {
session, err := c.cfg.RootClient.ExtendWebSession(ctx, authclient.WebSessionReq{
User: c.cfg.User,
PrevSessionID: c.cfg.Session.GetName(),
AccessRequestID: req.AccessRequestID,
Switchback: req.Switchback,
ReloadUser: req.ReloadUser,
})
if err != nil {
return nil, trace.Wrap(err)
}
return session, nil
}
// GetAgent returns agent that can be used to answer challenges
// for the web to ssh connection as well as certificate
func (c *SessionContext) GetAgent() (agent.ExtendedAgent, *ssh.Certificate, error) {
cert, err := c.GetSSHCertificate()
if err != nil {
return nil, nil, trace.Wrap(err)
}
if len(cert.ValidPrincipals) == 0 {
return nil, nil, trace.BadParameter("expected at least valid principal in certificate")
}
privateKey, err := ssh.ParseRawPrivateKey(c.cfg.Session.GetSSHPriv())
if err != nil {
return nil, nil, trace.Wrap(err, "failed to parse SSH private key")
}
keyring, ok := agent.NewKeyring().(agent.ExtendedAgent)
if !ok {
return nil, nil, trace.Errorf("unexpected keyring type: %T, expected agent.ExtendedKeyring", keyring)
}
err = keyring.Add(agent.AddedKey{
PrivateKey: privateKey,
Certificate: cert,
})
if err != nil {
return nil, nil, trace.Wrap(err)
}
return keyring, cert, nil
}
func (c *SessionContext) getCheckers() ([]ssh.PublicKey, error) {
ctx := context.TODO()
cas, err := c.cfg.UnsafeCachedAuthClient.GetCertAuthorities(ctx, types.HostCA, false)
if err != nil {
return nil, trace.Wrap(err)
}
var keys []ssh.PublicKey
for _, ca := range cas {
checkers, err := sshutils.GetCheckers(ca)
if err != nil {
return nil, trace.Wrap(err)
}
keys = append(keys, checkers...)
}
return keys, nil
}
// GetSSHCertificate returns the *ssh.Certificate associated with this session.
func (c *SessionContext) GetSSHCertificate() (*ssh.Certificate, error) {
return apisshutils.ParseCertificate(c.cfg.Session.GetPub())
}
// GetX509Certificate returns the *x509.Certificate associated with this session.
func (c *SessionContext) GetX509Certificate() (*x509.Certificate, error) {
tlsCert, err := tlsca.ParseCertificatePEM(c.cfg.Session.GetTLSCert())
if err != nil {
return nil, trace.Wrap(err)
}
return tlsCert, nil
}
// GetUserAccessChecker returns AccessChecker derived from the SSH certificate
// associated with this session.
func (c *SessionContext) GetUserAccessChecker() (services.AccessChecker, error) {
cert, err := c.GetSSHCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
ident, err := sshca.DecodeIdentity(cert)
if err != nil {
return nil, trace.Wrap(err)
}
accessInfo := services.AccessInfoFromLocalSSHIdentity(ident)
accessChecker, err := services.NewAccessCheckerForUserSession(accessInfo, c.cfg.RootClusterName, c.cfg.UnsafeCachedAuthClient)
return accessChecker, trace.Wrap(err)
}
func (c *SessionContext) GetUserScopedAccessCheckerContext(ctx context.Context) (*services.ScopedAccessCheckerContext, error) {
cert, err := c.GetSSHCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
ident, err := sshca.DecodeIdentity(cert)
if err != nil {
return nil, trace.Wrap(err)
}
accessInfo := services.AccessInfoFromLocalSSHIdentity(ident)
if accessInfo.ScopePin == nil {
checker, err := c.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
return services.NewScopedAccessCheckerContextFromUnscoped(checker), nil
}
checkerCtx, err := services.NewScopedAccessCheckerContext(
ctx, accessInfo, c.cfg.RootClusterName, c.cfg.UnsafeScopedRoleReader,
)
if err != nil {
return nil, trace.Wrap(err)
}
return checkerCtx, err
}
// GetProxyListenerMode returns cluster proxy listener mode form cluster networking config.
func (c *SessionContext) GetProxyListenerMode(ctx context.Context) (types.ProxyListenerMode, error) {
resp, err := c.cfg.UnsafeCachedAuthClient.GetClusterNetworkingConfig(ctx)
if err != nil {
return types.ProxyListenerMode_Separate, trace.Wrap(err)
}
return resp.GetProxyListenerMode(), nil
}
// GetIdentity returns identity parsed from the session's TLS certificate.
func (c *SessionContext) GetIdentity() (*tlsca.Identity, error) {
cert, err := c.GetX509Certificate()
if err != nil {
return nil, trace.Wrap(err)
}
identity, err := tlsca.FromSubject(cert.Subject, cert.NotAfter)
if err != nil {
return nil, trace.Wrap(err)
}
return identity, nil
}
// GetSessionID returns the ID of the underlying user web session.
func (c *SessionContext) GetSessionID() string {
return c.cfg.Session.GetName()
}
// GetRootClusterName returns the root cluster name.
func (c *SessionContext) GetRootClusterName() string {
return c.cfg.RootClusterName
}
// Close cleans up resources associated with this context and removes it
// from the user context
func (c *SessionContext) Close() error {
var err error
c.mu.Lock()
if c.kubeGRPCServiceConn != nil {
err = c.kubeGRPCServiceConn.Close()
}
c.mu.Unlock()
return trace.NewAggregate(c.remoteClientCache.Close(), c.cfg.RootClient.Close(), err)
}
// getToken returns the bearer token associated with the underlying
// session. Note that sessions are separate from bearer tokens and this
// is only useful immediately after a session has been created to query
// the token.
func (c *SessionContext) getToken() (types.WebToken, error) {
t, err := types.NewWebToken(c.cfg.Session.GetBearerTokenExpiryTime(), types.WebTokenSpecV3{
User: c.cfg.Session.GetUser(),
Token: c.cfg.Session.GetBearerToken(),
})
if err != nil {
return nil, trace.Wrap(err)
}
return t, nil
}
// expired returns whether this context has expired.
// Session records in the backend are created with a built-in expiry that
// automatically deletes the session record from the back-end database at
// the end of its natural life.
// If a session record still exists in the backend, it is considered still
// alive, regardless of the time. If no such record exists then a record is
// considered expired when its bearer token TTL is in the past (subject to
// lingering threshold)
func (c *SessionContext) expired(ctx context.Context) bool {
_, err := c.cfg.Parent.readSession(ctx, types.GetWebSessionRequest{
User: c.cfg.User,
SessionID: c.cfg.Session.GetName(),
})
switch {
case err == nil:
// If looking up the session in the cache or backend succeeds, then
// it by definition must not have expired yet.
return false
case trace.IsNotFound(err):
// If the session doesn't exist in the cache or backend, then it
// was removed during user logout, expire the session immediately.
return true
default:
c.cfg.Log.DebugContext(ctx, "Failed to query web session", "error", err)
}
expiry := c.cfg.Session.GetEarliestExpiry()
// If the session has no expiry time, then also by definition it
// cannot be expired
if expiry.IsZero() {
return false
}
// Give the session some time to linger so existing users of the context
// have successfully disposed of them.
// If we remove the session immediately, a stale copy might still use the
// cached cluster clients.
// This is a cheaper way to avoid race without introducing object
// reference counters.
return c.cfg.Parent.clock.Since(expiry) > c.cfg.Parent.sessionLingeringThreshold
}
// cachedSessionLingeringThreshold specifies the maximum amount of time the session cache
// will hold onto a session before removing it. This period allows all outstanding references
// to disappear without fear of racing with the removal
const cachedSessionLingeringThreshold = 2 * time.Minute
type sessionCacheOptions struct {
proxyClient authclient.ClientI
scopedRoleReader services.ScopedRoleReader
accessPoint authclient.ReadProxyAccessPoint
servers []utils.NetAddr
cipherSuites []uint16
clock clockwork.Clock
// sessionLingeringThreshold specifies the time the session will linger
// in the cache before getting purged after it has expired
sessionLingeringThreshold time.Duration
// proxySigner is used to sign PROXY header and securely propagate client's real IP
proxySigner multiplexer.PROXYHeaderSigner
// See [sessionCache.sessionWatcherStartImmediately]. Used for testing.
sessionWatcherStartImmediately bool
// See [sessionCache.sessionWatcherInitializedChannel]. Used for testing.
sessionWatcherInitializedChannel chan struct{}
// See [sessionCache.sessionWatcherEventProcessedChannel]. Used for testing.
sessionWatcherEventProcessedChannel chan struct{}
logger *slog.Logger
buildType string
// See [sessionCache.rootClientDialOptions]. Used for testing.
rootClientDialOptions []grpc.DialOption
}
// newSessionCache creates a [sessionCache] from the provided [config] and
// launches a goroutine that runs until [ctx] is completed which
// periodically purges invalid sessions.
func newSessionCache(ctx context.Context, config sessionCacheOptions) (*sessionCache, error) {
clusterName, err := config.proxyClient.GetClusterName(ctx)
if err != nil {
return nil, trace.Wrap(err)
}
if config.clock == nil {
config.clock = clockwork.NewRealClock()
}
if config.logger == nil {
config.logger = slog.Default()
}
cache := &sessionCache{
clusterName: clusterName.GetClusterName(),
proxyClient: config.proxyClient,
scopedRoleReader: config.scopedRoleReader,
accessPoint: config.accessPoint,
sessions: make(map[string]*SessionContext),
resources: make(map[string]*sessionResources),
authServers: config.servers,
closer: utils.NewCloseBroadcaster(),
cipherSuites: config.cipherSuites,
log: config.logger,
clock: config.clock,
sessionLingeringThreshold: config.sessionLingeringThreshold,
proxySigner: config.proxySigner,
sessionWatcherStartImmediately: config.sessionWatcherStartImmediately,
sessionWatcherInitializedChannel: config.sessionWatcherInitializedChannel,
sessionWatcherMarkInitialized: sync.OnceFunc(func() {
c := config.sessionWatcherInitializedChannel
if c != nil {
close(c)
}
}),
sessionWatcherEventProcessedChannel: config.sessionWatcherEventProcessedChannel,
buildType: config.buildType,
rootClientDialOptions: config.rootClientDialOptions,
}
// periodically close expired and unused sessions
go cache.expireSessions(ctx)
// Watch for session updates.
go cache.watchWebSessions(ctx)
return cache, nil
}
// sessionCache handles web session authentication,
// and holds in-memory contexts associated with each session
type sessionCache struct {
log *slog.Logger
proxyClient authclient.ClientI
authServers []utils.NetAddr
accessPoint authclient.ReadProxyAccessPoint
scopedRoleReader services.ScopedRoleReader
closer *utils.CloseBroadcaster
clusterName string
clock clockwork.Clock
// sessionLingeringThreshold specifies the time the session will linger
// in the cache before getting purged after it has expired
sessionLingeringThreshold time.Duration
// cipherSuites is the list of supported TLS cipher suites.
cipherSuites []uint16
buildType string
mu sync.RWMutex
// sessions maps user/sessionID to an active web session value between renewals.
// This is the client-facing session handle
sessions map[string]*SessionContext
// sessionGroup ensures only a single SessionContext will exist for a
// user+session.
sessionGroup singleflight.Group
// session cache maintains a list of resources per-user as long
// as the user session is active even though individual session values
// are periodically recycled.
// Resources are disposed of when the corresponding session
// is either explicitly invalidated (e.g. during logout) or the
// resources are themselves closing
resources map[string]*sessionResources
// proxySigner is used to sign PROXY header and securely propagate client's real IP
proxySigner multiplexer.PROXYHeaderSigner
// sessionWatcherStartImmediately removes the First component of the linear
// backoff used to start the WebSession watcher.
// Used for testing.
sessionWatcherStartImmediately bool
// sessionWatcherInitializedChannel is used to signal that the sessionWatcher
// received its first OpInit event and is ready to observe updates.
// May be nil.
// Used for testing.
sessionWatcherInitializedChannel chan struct{}
// sessionWatcherMarkInitialized safely closes
// sessionWatcherInitializedChannel.
sessionWatcherMarkInitialized func()
// sessionWatcherEventProcessedChannel is used to signal that the
// sessionWatcher processed an event.
// May be nil.
// Used for testing.
sessionWatcherEventProcessedChannel chan struct{}
// rootClientDialOptions contains additional gRPC dial options for root clients.
// Used for testing.
rootClientDialOptions []grpc.DialOption
}
// Close closes all allocated resources and stops goroutines
func (s *sessionCache) Close() error {
s.log.InfoContext(context.Background(), "Closing session cache")
return s.closer.Close()
}
func (s *sessionCache) ActiveSessions() int {
s.mu.RLock()
defer s.mu.RUnlock()
return len(s.sessions)
}
func (s *sessionCache) expireSessions(ctx context.Context) {
ticker := s.clock.NewTicker(1 * time.Second)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.Chan():
s.clearExpiredSessions(ctx)
case <-s.closer.C:
return
}
}
}
func (s *sessionCache) clearExpiredSessions(ctx context.Context) {
s.mu.Lock()
defer s.mu.Unlock()
for _, c := range s.sessions {
if !c.expired(ctx) {
continue
}
s.removeSessionContextLocked(ctx, c.cfg.Session.GetUser(), c.cfg.Session.GetName())
s.log.DebugContext(ctx, "Context expired", "context", logutils.StringerAttr(c))
}
}
// watchWebSessions runs the WebSession watcher loop.
// It only stops when ctx is done.
func (s *sessionCache) watchWebSessions(ctx context.Context) {
// Watcher not necessary for OSS.
if s.buildType != modules.BuildEnterprise {
return
}
linear := utils.NewDefaultLinear(s.clock)
if s.sessionWatcherStartImmediately {
linear.First = 0
}
s.log.DebugContext(ctx, "sessionCache: Starting WebSession watcher")
for {
select {
// Stop when the context tells us to.
case <-ctx.Done():
s.log.DebugContext(ctx, "sessionCache: Stopping WebSession watcher")
return
case <-linear.After():
linear.Inc()
}
if err := s.watchWebSessionsOnce(ctx, linear.Reset); err != nil && !errors.Is(err, context.Canceled) {
const msg = "" +
"sessionCache: WebSession watcher aborted, re-connecting. " +
"This may have an impact on Device Trust web sessions."
s.log.WarnContext(ctx, msg, "error", err)
}
}
}
// watchWebSessionsOnce creates a watcher for WebSessions and watches for its
// events.
//
// Any session updated with device extensions is evicted from the cache. That is
// so the new certificates are forcefully loaded by the Proxy.
//
// Sessions updated for other reasons (no device extensions present) or cached
// sessions that already have device extensions are ignored. This avoids
// disconnecting clients during periodic bearer token refresh by the Web UI.
func (s *sessionCache) watchWebSessionsOnce(ctx context.Context, reset func()) error {
watcher, err := s.proxyClient.NewWatcher(ctx, types.Watch{
Name: teleport.ComponentWebProxy + ".sessionCache." + types.KindWebSession,
Kinds: []types.WatchKind{
{
Kind: types.KindWebSession,
// Watch only for KindWebSession.
// SubKinds include KindAppSession, KindSnowflakeSession, etc.
SubKind: types.KindWebSession,
},
},
})
if err != nil {
return trace.Wrap(err)
}
defer watcher.Close()
// notifyProcessed is a feedback mechanism for tests.
notifyProcessed := func() {
if s.sessionWatcherEventProcessedChannel != nil {
s.sessionWatcherEventProcessedChannel <- struct{}{}
}
}
for {
select {
case <-ctx.Done():
return ctx.Err()
case <-watcher.Done():
return errors.New("watcher closed")
case event := <-watcher.Events():
reset() // Reset linear backoff attempts.
s.log.Log(ctx, logutils.TraceLevel, "sessionCache: Received watcher event",
"event", logutils.StringerAttr(event),
)
if event.Type == types.OpInit {
s.sessionWatcherMarkInitialized()
continue
}
if event.Type != types.OpPut {
continue // We only care about OpPut at the moment.
}
session, ok := event.Resource.(types.WebSession)
if !ok {
s.log.WarnContext(ctx, "sessionCache: Received unexpected resource type",
"resource_type", logutils.TypeAttr(event.Resource),
)
continue
}
if !session.GetHasDeviceExtensions() {
s.log.DebugContext(ctx, "sessionCache: Updated session doesn't have device extensions, skipping",
"session_id", session.GetName(),
)
notifyProcessed()
continue
}
// Release existing and non-device-aware session.
if err := s.releaseResourcesIfNoDeviceExtensions(ctx, session.GetUser(), session.GetName()); err != nil {
s.log.DebugContext(ctx, "sessionCache: Failed to release updated session",
"error", err,
"session_id", session.GetName(),
)
}
notifyProcessed()
}
}
}
func (s *sessionCache) releaseResourcesIfNoDeviceExtensions(ctx context.Context, user, sessionID string) error {
s.mu.Lock()
defer s.mu.Unlock()
id := sessionKey(user, sessionID)
switch sessionCtx, ok := s.sessions[id]; {
case !ok:
return nil // Session not found
case sessionCtx.cfg.Session.GetHasDeviceExtensions():
s.log.DebugContext(ctx, "sessionCache: Session already has device extensions, skipping",
"session_id", sessionID,
)
return nil
}
s.log.DebugContext(ctx, "sessionCache: Releasing session resources due to device extensions upgrade",
"session_id", sessionID,
)
return s.releaseResourcesLocked(ctx, user, sessionID)
}
// AuthWithOTP authenticates the specified user with the given password and OTP token.
// Returns a new web session if successful.
func (s *sessionCache) AuthWithOTP(
ctx context.Context,
user, pass, otpToken string,
scope string,
clientMeta *authclient.ForwardedClientMetadata,
) (types.WebSession, error) {
return s.proxyClient.AuthenticateWebUser(ctx, authclient.AuthenticateUserRequest{
Username: user,
Pass: &authclient.PassCreds{Password: []byte(pass)},
OTP: &authclient.OTPCreds{
Password: []byte(pass),
Token: otpToken,
},
ClientMetadata: clientMeta,
Scope: scope,
})
}
// AuthWithoutOTP authenticates the specified user with the given password.
// Returns a new web session if successful.
func (s *sessionCache) AuthWithoutOTP(
ctx context.Context, user, pass string, scope string, clientMeta *authclient.ForwardedClientMetadata,
) (types.WebSession, error) {
return s.proxyClient.AuthenticateWebUser(ctx, authclient.AuthenticateUserRequest{
Username: user,
Pass: &authclient.PassCreds{
Password: []byte(pass),
},
ClientMetadata: clientMeta,
Scope: scope,
})
}
func (s *sessionCache) AuthenticateWebUser(
ctx context.Context, req *client.AuthenticateWebUserRequest, clientMeta *authclient.ForwardedClientMetadata,
) (types.WebSession, error) {
authReq := authclient.AuthenticateUserRequest{
Username: req.User,
ClientMetadata: clientMeta,
Scope: req.Scope,
}
if req.WebauthnAssertionResponse != nil {
authReq.Webauthn = req.WebauthnAssertionResponse
}
return s.proxyClient.AuthenticateWebUser(ctx, authReq)
}
func (s *sessionCache) AuthenticateSSHUser(
ctx context.Context, c client.AuthenticateSSHUserRequest, clientMeta *authclient.ForwardedClientMetadata,
) (*authclient.CLILoginResponse, error) {
authReq := authclient.AuthenticateUserRequest{
Username: c.User,
Scope: c.Scope,
ClientMetadata: clientMeta,
SSHPublicKey: c.UserPublicKeys.SSHPubKey,
TLSPublicKey: c.UserPublicKeys.TLSPubKey,
}
if c.Password != "" {
authReq.Pass = &authclient.PassCreds{Password: []byte(c.Password)}
}
if c.WebauthnChallengeResponse != nil {
authReq.Webauthn = c.WebauthnChallengeResponse
}
if c.TOTPCode != "" {
authReq.OTP = &authclient.OTPCreds{
Password: []byte(c.Password),
Token: c.TOTPCode,
}
}
if c.BrowserMFAResponse != nil {
authReq.BrowserMFA = &proto.BrowserMFAResponse{
RequestId: c.BrowserMFAResponse.RequestID,
WebauthnResponse: webauthntypes.CredentialAssertionResponseToProto(c.BrowserMFAResponse.WebauthnResponse),
}
}
return s.proxyClient.AuthenticateSSHUser(ctx, authclient.AuthenticateSSHRequest{
AuthenticateUserRequest: authReq,
CompatibilityMode: c.Compatibility,
TTL: c.TTL,
RouteToCluster: c.RouteToCluster,
KubernetesCluster: c.KubernetesCluster,
SSHAttestationStatement: c.UserPublicKeys.SSHAttestationStatement,
TLSAttestationStatement: c.UserPublicKeys.TLSAttestationStatement,
})
}
// Ping gets basic info about the auth server.
func (s *sessionCache) Ping(ctx context.Context) (proto.PingResponse, error) {
return s.proxyClient.Ping(ctx)
}
func (s *sessionCache) ValidateTrustedCluster(ctx context.Context, validateRequest *authclient.ValidateTrustedClusterRequest) (*authclient.ValidateTrustedClusterResponse, error) {
return s.proxyClient.ValidateTrustedCluster(ctx, validateRequest)
}
// getOrCreateSession gets the SessionContext for the user and session ID. If one does
// not exist, then a new one is created.
func (s *sessionCache) getOrCreateSession(ctx context.Context, user, sessionID string) (*SessionContext, error) {
key := sessionKey(user, sessionID)
// Use sessionGroup to prevent multiple requests from racing to create a SessionContext.
i, err, _ := s.sessionGroup.Do(key, func() (any, error) {
sessionCtx, ok := s.getContext(key)
if ok {
return sessionCtx, nil
}
return s.newSessionContext(ctx, user, sessionID)
})
if err != nil {
return nil, trace.Wrap(err)
}
sctx, ok := i.(*SessionContext)
if !ok {
return nil, trace.BadParameter("expected SessionContext, got %T", i)
}
identity, err := sctx.GetIdentity()
if err != nil {
return nil, trace.Wrap(err)
}
// Enforce IP Pinning if it is present in the user's certificate.
var clientAddr string
if clientSrcAddr, err := authz.ClientSrcAddrFromContext(ctx); err == nil {
clientAddr = clientSrcAddr.String()
}
if err := authz.CheckIPPinning(ctx, clientAddr, identity.PinnedIP, false, s.log); err != nil {
return nil, trace.Wrap(err)
}
return sctx, nil
}
func (s *sessionCache) invalidateSession(ctx context.Context, sctx *SessionContext) error {
defer sctx.Close()
clt, err := sctx.GetClient()
if err != nil {
return trace.Wrap(err)
}
// App session, SAML session and web session deletion should be treated as a single transaction.
// To avoid aborting deletion midpoint due to a failure in one of the session deletion,
// we use sessionDeletionErr below to join errors and return them at last.
var sessionDeletionErrs error
if err := clt.DeleteUserAppSessions(ctx, &proto.DeleteUserAppSessionsRequest{Username: sctx.GetUser()}); err != nil {
sessionDeletionErrs = err
}
// Delete just the session - leave the bearer token to linger to avoid
// failing a client query still using the old token.
if err := clt.WebSessions().Delete(ctx, types.DeleteWebSessionRequest{
User: sctx.GetUser(),
SessionID: sctx.GetSessionID(),
}); err != nil && !trace.IsNotFound(err) {
sessionDeletionErrs = errors.Join(sessionDeletionErrs, err)
}
return trace.Wrap(sessionDeletionErrs)
}
func (s *sessionCache) getContext(key string) (*SessionContext, bool) {
s.mu.RLock()
defer s.mu.RUnlock()
ctx, ok := s.sessions[key]
return ctx, ok
}
func (s *sessionCache) insertContext(user string, sctx *SessionContext) (exists bool) {
s.mu.Lock()
defer s.mu.Unlock()
id := sessionKey(user, sctx.GetSessionID())
if _, exists := s.sessions[id]; exists {
return true
}
s.sessions[id] = sctx
return false
}
func (s *sessionCache) releaseResources(ctx context.Context, user, sessionID string) error {
s.mu.Lock()
defer s.mu.Unlock()
return s.releaseResourcesLocked(ctx, user, sessionID)
}
func (s *sessionCache) removeSessionContextLocked(ctx context.Context, user, sessionID string) error {
id := sessionKey(user, sessionID)
sess, ok := s.sessions[id]
if !ok {
return nil
}
delete(s.sessions, id)
err := sess.Close()
if err != nil {
s.log.WarnContext(ctx, "Failed to close session context",
"context", logutils.StringerAttr(sess),
"error", err,
)
return trace.Wrap(err)
}
return nil
}
func (s *sessionCache) releaseResourcesLocked(ctx context.Context, user, sessionID string) error {
var errors []error
err := s.removeSessionContextLocked(ctx, user, sessionID)
if err != nil {
errors = append(errors, err)
}
if sess, ok := s.resources[user]; ok {
delete(s.resources, user)
if err := sess.Close(); err != nil {
s.log.WarnContext(ctx, "Failed to clean up session context", "error", err)
errors = append(errors, err)
}
}
return trace.NewAggregate(errors...)
}
func (s *sessionCache) upsertSessionContext(user string) *sessionResources {
s.mu.Lock()
defer s.mu.Unlock()
if ctx, exists := s.resources[user]; exists {
return ctx
}
ctx := &sessionResources{
log: s.log.With(
teleport.ComponentKey, "user-session",
"user", user,
),
}
s.resources[user] = ctx
return ctx
}
// newSessionContext creates a new web session context for the specified user/session ID
func (s *sessionCache) newSessionContext(ctx context.Context, user, sessionID string) (*SessionContext, error) {
session, err := s.proxyClient.AuthenticateWebUser(ctx, authclient.AuthenticateUserRequest{
Username: user,
Session: &authclient.SessionCreds{
ID: sessionID,
},
})
if err != nil {
// This will fail if the session has expired and was removed
return nil, trace.Wrap(err)
}
return s.newSessionContextFromSession(ctx, session)
}
func (s *sessionCache) newSessionContextFromSession(ctx context.Context, session types.WebSession) (*SessionContext, error) {
tlsConfig, err := s.tlsConfig(ctx, session.GetTLSCert(), session.GetTLSPriv())
if err != nil {
return nil, trace.Wrap(err)
}
// Enforce IP Pinning if it is present in the user's certificate.
cert, err := tlsca.ParseCertificatePEM(session.GetTLSCert())
if err != nil {
return nil, trace.Wrap(err)
}
identity, err := tlsca.FromSubject(cert.Subject, cert.NotAfter)
if err != nil {
return nil, trace.Wrap(err)
}
var clientAddr string
if clientSrcAddr, err := authz.ClientSrcAddrFromContext(ctx); err == nil {
clientAddr = clientSrcAddr.String()
}
if err := authz.CheckIPPinning(ctx, clientAddr, identity.PinnedIP, false, s.log); err != nil {
return nil, trace.Wrap(err)
}
userClient, err := authclient.NewClient(apiclient.Config{
Addrs: utils.NetAddrsToStrings(s.authServers),
Credentials: []apiclient.Credentials{apiclient.LoadTLS(tlsConfig)},
CircuitBreakerConfig: breaker.NoopBreakerConfig(),
PROXYHeaderGetter: client.CreatePROXYHeaderGetter(ctx, s.proxySigner),
DialOpts: s.rootClientDialOptions,
})
if err != nil {
return nil, trace.Wrap(err)
}
sctx, err := NewSessionContext(SessionContextConfig{
Log: s.log.With(
"user", session.GetUser(),
"session", session.GetShortName(),
),
User: session.GetUser(),
RootClient: userClient,
UnsafeCachedAuthClient: s.accessPoint,
UnsafeScopedRoleReader: s.scopedRoleReader,
Parent: s,
Resources: s.upsertSessionContext(session.GetUser()),
Session: session,
RootClusterName: s.clusterName,
})
if err != nil {
return nil, trace.Wrap(err)
}
if exists := s.insertContext(session.GetUser(), sctx); exists {
// this means that someone has just inserted the context, so
// close our extra context and return
sctx.Close()
}
return sctx, nil
}
func (s *sessionCache) tlsConfig(ctx context.Context, cert, privKey []byte) (*tls.Config, error) {
ca, err := s.proxyClient.GetCertAuthority(ctx, types.CertAuthID{
Type: types.HostCA,
DomainName: s.clusterName,
}, false)
if err != nil {
return nil, trace.Wrap(err)
}
certPool, err := services.CertPool(ca)
if err != nil {
return nil, trace.Wrap(err)
}
tlsConfig := utils.TLSConfig(s.cipherSuites)
tlsCert, err := tls.X509KeyPair(cert, privKey)
if err != nil {
return nil, trace.Wrap(err, "failed to parse TLS certificate and key")
}
tlsConfig.Certificates = []tls.Certificate{tlsCert}
tlsConfig.RootCAs = certPool
tlsConfig.ServerName = apiutils.EncodeClusterName(s.clusterName)
tlsConfig.Time = s.clock.Now
return tlsConfig, nil
}
func (s *sessionCache) readSession(ctx context.Context, req types.GetWebSessionRequest) (types.WebSession, error) {
// Read session from the cache first
session, err := s.accessPoint.GetWebSession(ctx, req)
if err == nil {
return session, nil
}
// Fallback to proxy otherwise
return s.proxyClient.GetWebSession(ctx, req)
}
func (s *sessionCache) readBearerToken(ctx context.Context, req types.GetWebTokenRequest) (types.WebToken, error) {
// Read token from the cache first
token, err := s.accessPoint.GetWebToken(ctx, req)
if err == nil {
return token, nil
}
// Fallback to proxy otherwise
return s.proxyClient.GetWebToken(ctx, req)
}
// Close releases all underlying resources for the user session.
func (c *sessionResources) Close() error {
closers := c.transferClosers()
var errors []error
for _, closer := range closers {
if err := closer.Close(); err != nil {
errors = append(errors, err)
}
}
return trace.NewAggregate(errors...)
}
// sessionResources persists resources initiated by a web session
// but which might outlive the session.
type sessionResources struct {
log *slog.Logger
mu sync.Mutex
closers []io.Closer
}
// addClosers adds the specified closers to this context
func (c *sessionResources) addClosers(closers ...io.Closer) {
c.mu.Lock()
defer c.mu.Unlock()
c.closers = append(c.closers, closers...)
}
// removeCloser removes the specified closer from this context
func (c *sessionResources) removeCloser(closer io.Closer) {
c.mu.Lock()
defer c.mu.Unlock()
for i, cls := range c.closers {
if cls == closer {
c.closers = slices.Delete(c.closers, i, i+1)
return
}
}
}
func (c *sessionResources) transferClosers() []io.Closer {
c.mu.Lock()
defer c.mu.Unlock()
closers := c.closers
c.closers = nil
return closers
}
func sessionKey(user, sessionID string) string {
return user + sessionID
}
// remoteClientCache stores remote clients keyed by cluster name while also keeping
// track of the actual remote cluster associated with the client (in case the
// remote cluster has changed). Safe for concurrent access. Closes all clients and
// wipes the cache on Close.
type remoteClientCache struct {
sync.Mutex
clients map[string]struct {
authclient.ClientI
reversetunnelclient.Cluster
}
}
func (c *remoteClientCache) addRemoteClient(cluster reversetunnelclient.Cluster, remoteClient authclient.ClientI) error {
c.Lock()
defer c.Unlock()
if c.clients == nil {
c.clients = make(map[string]struct {
authclient.ClientI
reversetunnelclient.Cluster
})
}
var err error
if c.clients[cluster.GetName()].ClientI != nil {
err = c.clients[cluster.GetName()].ClientI.Close()
}
c.clients[cluster.GetName()] = struct {
authclient.ClientI
reversetunnelclient.Cluster
}{remoteClient, cluster}
return err
}
func (c *remoteClientCache) getRemoteClient(cluster reversetunnelclient.Cluster) (authclient.ClientI, bool) {
c.Lock()
defer c.Unlock()
remoteClt, ok := c.clients[cluster.GetName()]
return remoteClt.ClientI, ok && remoteClt.Cluster == cluster
}
func (c *remoteClientCache) Close() error {
c.Lock()
defer c.Unlock()
errors := make([]error, 0, len(c.clients))
for _, clt := range c.clients {
errors = append(errors, clt.ClientI.Close())
}
c.clients = nil
return trace.NewAggregate(errors...)
}
/*
* 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 (
"bytes"
"fmt"
"net/http"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/client/db"
"github.com/gravitational/teleport/lib/client/identityfile"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/utils"
)
/*
signDatabaseCertificate returns the necessary files to set up mTLS using the `db` format
This is the equivalent of running the tctl command
As an example, requesting:
POST /webapi/sites/mycluster/sign/db
{
"hostname": "pg.example.com",
"ttl": "2190h"
}
Should be equivalent to running:
tctl auth sign --host=pg.example.com --ttl=2190h --format=db
This endpoint returns a tar.gz compressed archive containing the required files to setup mTLS for the database.
*/
func (h *Handler) signDatabaseCertificate(w http.ResponseWriter, r *http.Request, p httprouter.Params, cluster reversetunnelclient.Cluster, token types.ProvisionToken) (any, error) {
if !token.GetRoles().Include(types.RoleDatabase) {
return nil, trace.AccessDenied("required '%s' role was not provided by the token", types.RoleDatabase)
}
req := &signDatabaseCertificateReq{}
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
virtualFS := identityfile.NewInMemoryConfigWriter()
dbCertReq := db.GenerateDatabaseCertificatesRequest{
ClusterAPI: h.auth.proxyClient,
Principals: []string{req.Hostname},
OutputFormat: identityfile.FormatDatabase,
OutputCanOverwrite: true,
OutputLocation: "server",
IdentityFileWriter: virtualFS,
TTL: req.TTL,
}
filesWritten, err := db.GenerateDatabaseServerCertificates(r.Context(), dbCertReq)
if err != nil {
return nil, trace.Wrap(err)
}
archiveName := fmt.Sprintf("teleport_mTLS_%s.tar.gz", req.Hostname)
archiveBytes, err := utils.CompressTarGzArchive(filesWritten, virtualFS)
if err != nil {
return nil, trace.Wrap(err)
}
// Set file name
w.Header().Set("Content-Disposition", fmt.Sprintf(`attachment;filename="%v"`, archiveName))
// ServeContent sets the correct headers: Content-Type, Content-Length and Accept-Ranges.
// It also handles the Range negotiation
http.ServeContent(w, r, archiveName, time.Now(), bytes.NewReader(archiveBytes.Bytes()))
return nil, nil
}
type signDatabaseCertificateReq struct {
Hostname string `json:"hostname,omitempty"`
TTLRaw string `json:"ttl,omitempty"`
TTL time.Duration `json:"-"`
}
// CheckAndSetDefaults will validate and convert the received values
// Hostname must not be empty
// TTL must either be a valid time.Duration or empty (inherits apidefaults.CertDuration)
func (s *signDatabaseCertificateReq) CheckAndSetDefaults() error {
if s.Hostname == "" {
return trace.BadParameter("missing hostname")
}
if s.TTLRaw == "" {
s.TTLRaw = apidefaults.CertDuration.String()
}
ttl, err := time.ParseDuration(s.TTLRaw)
if err != nil {
return trace.BadParameter("invalid ttl '%s', use https://pkg.go.dev/time#ParseDuration format (example: 2190h)", s.TTLRaw)
}
s.TTL = ttl
return nil
}
// Teleport
// Copyright (C) 2024 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 (
"net/http"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/spiffe/go-spiffe/v2/bundle/spiffebundle"
"github.com/spiffe/go-spiffe/v2/spiffeid"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/lib/jwt"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/tlsca"
"github.com/gravitational/teleport/lib/utils/oidc"
)
// getSPIFFEBundle returns the SPIFFE-compatible trust bundle which allows other
// trust domains to federate with this Teleport cluster.
//
// Mounted at /webapi/spiffe/bundle.json
//
// Must abide by the standard for a "https_web" profile as described in
// https://github.com/spiffe/spiffe/blob/main/standards/SPIFFE_Federation.md#5-serving-and-consuming-a-spiffe-bundle-endpoint
func (h *Handler) getSPIFFEBundle(w http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
cn, err := h.GetAccessPoint().GetClusterName(r.Context())
if err != nil {
return nil, trace.Wrap(err, "fetching cluster name")
}
td, err := spiffeid.TrustDomainFromString(cn.GetClusterName())
if err != nil {
return nil, trace.Wrap(err, "creating trust domain")
}
bundle := spiffebundle.New(td)
// The refresh hint indicates how often a federated trust domain should
// check for updates to the bundle. This should be a low value to ensure
// that CA rotations are picked up quickly. Since we're leveraging
// https_web, it's not critical for a federated trust domain to catch
// all phases of the rotation - however, if we support https_spiffe in
// future, we may need to consider a lower value or enforcing a wait
// period during rotations equivalent to the refresh hint.
bundle.SetRefreshHint(5 * time.Minute)
// TODO(noah):
// For now, we omit the SequenceNumber field. This is only a SHOULD not a
// MUST per the spec. To add this, we will add a sequence number to the
// cert authority and increment it on every update.
const loadKeysFalse = false
spiffeCA, err := h.GetAccessPoint().GetCertAuthority(r.Context(), types.CertAuthID{
Type: types.SPIFFECA,
DomainName: cn.GetClusterName(),
}, loadKeysFalse)
if err != nil {
return nil, trace.Wrap(err, "fetching SPIFFE CA")
}
// Add X509 authorities to the trust bundle.
for _, certPEM := range services.GetTLSCerts(spiffeCA) {
cert, err := tlsca.ParseCertificatePEM(certPEM)
if err != nil {
return nil, trace.Wrap(err, "parsing certificate")
}
bundle.AddX509Authority(cert)
}
// Add JWT authorities to the trust bundle.
for _, keyPair := range spiffeCA.GetTrustedJWTKeyPairs() {
pubKey, err := keys.ParsePublicKey(keyPair.PublicKey)
if err != nil {
return nil, trace.Wrap(err, "parsing public key")
}
kid, err := jwt.KeyID(pubKey)
if err != nil {
return nil, trace.Wrap(err, "generating key ID")
}
if err := bundle.AddJWTAuthority(kid, pubKey); err != nil {
return nil, trace.Wrap(err, "adding JWT authority to bundle")
}
}
bundleBytes, err := bundle.Marshal()
if err != nil {
return nil, trace.Wrap(err, "marshaling bundle")
}
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusOK)
if _, err = w.Write(bundleBytes); err != nil {
h.logger.DebugContext(h.cfg.Context, "Failed to write SPIFFE bundle response", "error", err)
}
return nil, nil
}
// Mounted at /workload-identity/.well-known/openid-configuration
func (h *Handler) getSPIFFEOIDCDiscoveryDocument(_ http.ResponseWriter, _ *http.Request, _ httprouter.Params) (any, error) {
issuer, err := oidc.IssuerFromPublicAddress(h.cfg.PublicProxyAddr, "/workload-identity")
if err != nil {
return nil, trace.Wrap(err, "determining issuer from public address")
}
return &oidc.OpenIDConfiguration{
Issuer: issuer,
JWKSURI: issuer + "/jwt-jwks.json",
Claims: []string{
"iss",
"sub",
"jti",
"aud",
"exp",
"iat",
},
IdTokenSigningAlgValuesSupported: []string{
"RS256",
},
ResponseTypesSupported: []string{
"id_token",
},
// Whilst this field is not required for GCP's Workload Identity
// Federation, it is required for AWS's AssumeRoleWithWebIdentity.
SubjectTypesSupported: []string{
"public",
},
}, nil
}
// Mounted at /workload-identity/jwt-jwks.json
func (h *Handler) getSPIFFEJWKS(_ http.ResponseWriter, r *http.Request, _ httprouter.Params) (any, error) {
return h.jwks(r.Context(), types.SPIFFECA, false)
}
/*
* 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"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"sync"
"sync/atomic"
"time"
"github.com/gogo/protobuf/proto"
"github.com/google/uuid"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/jonboulle/clockwork"
oteltrace "go.opentelemetry.io/otel/trace"
"golang.org/x/crypto/ssh"
"github.com/gravitational/teleport"
authproto "github.com/gravitational/teleport/api/client/proto"
mfav2 "github.com/gravitational/teleport/api/gen/proto/go/teleport/mfa/v2"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
"github.com/gravitational/teleport/api/mfa"
"github.com/gravitational/teleport/api/observability/tracing"
tracessh "github.com/gravitational/teleport/api/observability/tracing/ssh"
apissh "github.com/gravitational/teleport/api/ssh"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/keys"
"github.com/gravitational/teleport/api/utils/sshutils"
"github.com/gravitational/teleport/lib/agentless"
"github.com/gravitational/teleport/lib/auth/authclient"
wantypes "github.com/gravitational/teleport/lib/auth/webauthntypes"
"github.com/gravitational/teleport/lib/client"
clientssh "github.com/gravitational/teleport/lib/client/ssh"
"github.com/gravitational/teleport/lib/client/sso"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/multiplexer"
"github.com/gravitational/teleport/lib/proxy"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/sshagent"
"github.com/gravitational/teleport/lib/sshca"
"github.com/gravitational/teleport/lib/utils"
"github.com/gravitational/teleport/lib/utils/diagnostics/latency"
"github.com/gravitational/teleport/lib/web/terminal"
)
// TerminalRequest describes a request to create a web-based terminal
// to a remote SSH server.
type TerminalRequest struct {
// Server describes a server to connect to (serverId|hostname[:port]).
Server string `json:"server_id"`
// Login is Linux username to connect as.
Login string `json:"login"`
// Term is the initial PTY size.
Term session.TerminalParams `json:"term"`
// JoinSessionID is a Teleport session ID to join as.
JoinSessionID session.ID `json:"sid"`
// ProxyHostPort is the address of the server to connect to.
ProxyHostPort string `json:"-"`
// InteractiveCommand is a command to execute
InteractiveCommand []string `json:"-"`
// KeepAliveInterval is the interval for sending ping frames to web client.
// This value is pulled from the cluster network config and
// guaranteed to be set to a nonzero value as it's enforced by the configuration.
KeepAliveInterval time.Duration
// ParticipantMode is the mode that determines what you can do when you join an active session.
ParticipantMode types.SessionParticipantMode `json:"mode"`
}
// UserAuthClient is a subset of the Auth API that performs
// operations on behalf of the user so that the correct RBAC is applied.
type UserAuthClient interface {
GetSessionTracker(ctx context.Context, sessionID string) (types.SessionTracker, error)
IsMFARequired(ctx context.Context, req *authproto.IsMFARequiredRequest) (*authproto.IsMFARequiredResponse, error)
CreateAuthenticateChallenge(ctx context.Context, req *authproto.CreateAuthenticateChallengeRequest) (*authproto.MFAAuthenticateChallenge, error)
GenerateUserCerts(ctx context.Context, req authproto.UserCertsRequest) (*authproto.Certs, error)
MaintainSessionPresence(ctx context.Context) (authproto.AuthService_MaintainSessionPresenceClient, error)
ListUnifiedResources(ctx context.Context, req *authproto.ListUnifiedResourcesRequest) (*authproto.ListUnifiedResourcesResponse, error)
MFAServiceClientV2() mfav2.MFAServiceClient
}
// NewTerminal creates a web-based terminal based on WebSockets and returns a
// new TerminalHandler.
func NewTerminal(ctx context.Context, cfg TerminalHandlerConfig) (*TerminalHandler, error) {
err := cfg.CheckAndSetDefaults()
if err != nil {
return nil, trace.Wrap(err)
}
_, span := cfg.tracer.Start(ctx, "NewTerminal")
defer span.End()
return &TerminalHandler{
sshBaseHandler: sshBaseHandler{
logger: cfg.Logger.With(
teleport.ComponentKey, teleport.ComponentWebsocket,
"session_id", cfg.SessionData.ID.String(),
),
ctx: cfg.SessionCtx,
userAuthClient: cfg.UserAuthClient,
localAccessPoint: cfg.LocalAccessPoint,
sessionData: cfg.SessionData,
keepAliveInterval: cfg.KeepAliveInterval,
proxyHostPort: cfg.ProxyHostPort,
proxyPublicAddr: cfg.ProxyPublicAddr,
interactiveCommand: cfg.InteractiveCommand,
router: cfg.Router,
tracer: cfg.tracer,
resolver: cfg.HostNameResolver,
sshDialTimeout: cfg.SSHDialTimeout,
fipsBuild: cfg.FIPSBuild,
},
displayLogin: cfg.DisplayLogin,
term: cfg.Term,
proxySigner: cfg.PROXYSigner,
participantMode: cfg.ParticipantMode,
tracker: cfg.Tracker,
presenceChecker: cfg.PresenceChecker,
websocketConn: cfg.WebsocketConn,
}, nil
}
// TerminalHandlerConfig contains the configuration options necessary to
// correctly set up the TerminalHandler
type TerminalHandlerConfig struct {
// Logger specifies the logger.
Logger *slog.Logger
// Term is the initial PTY size.
Term session.TerminalParams
// SessionCtx is the context for the users web session.
SessionCtx *SessionContext
// UserAuthClient is used to fetch nodes and sessions from the backend.
UserAuthClient UserAuthClient
// LocalAccessPoint is the subset of the Proxy cache required to
// look up information from the local cluster. This should not
// be used for anything that requires RBAC on behalf of the user.
// Requests that should be made on behalf of the user should
// use [UserAuthClient].
LocalAccessPoint localAccessPoint
// HostNameResolver allows the hostname to be determined from a server UUID
// so that a friendly name can be displayed in the console tab.
HostNameResolver func(serverID string) (hostname string, err error)
// DisplayLogin is the login name to display in the UI.
DisplayLogin string
// SessionData is the data to send to the client on the initial session creation.
SessionData session.Session
// KeepAliveInterval is the interval for sending ping frames to web client.
// This value is pulled from the cluster network config and
// guaranteed to be set to a nonzero value as it's enforced by the configuration.
KeepAliveInterval time.Duration
// ProxyHostPort is the address of the server to connect to.
ProxyHostPort string
// ProxyPublicAddr is the public web proxy address.
ProxyPublicAddr string
// InteractiveCommand is a command to execute.
InteractiveCommand []string
// Router determines how connections to nodes are created
Router *proxy.Router
// TracerProvider is used to create the tracer
TracerProvider oteltrace.TracerProvider
// PROXYSigner is used to sign PROXY header and securely propagate client IP information
PROXYSigner multiplexer.PROXYHeaderSigner
// tracer is used to create spans
tracer oteltrace.Tracer
// ParticipantMode is the mode that determines what you can do when you join an active session.
ParticipantMode types.SessionParticipantMode
// Tracker is the session tracker of the session being joined. May be nil
// if the user is not joining a session.
Tracker types.SessionTracker
// PresenceChecker used for presence checking.
PresenceChecker PresenceChecker
// Clock allows interaction with time.
Clock clockwork.Clock
// WebsocketConn is the active websocket connection
WebsocketConn *websocket.Conn
// SSHDialTimeout is the dial timeout that should be enforced on ssh connections.
SSHDialTimeout time.Duration
// FIPSBuild indicates if this is a Teleport FIPS build.
FIPSBuild bool
}
func (t *TerminalHandlerConfig) CheckAndSetDefaults() error {
if t.Logger == nil {
t.Logger = slog.Default().With(teleport.ComponentKey, teleport.ComponentWebsocket)
}
// Make sure whatever session is requested is a valid session id.
if !t.SessionData.ID.IsZero() {
_, err := session.ParseID(t.SessionData.ID.String())
if err != nil {
return trace.BadParameter("sid: invalid session id")
}
}
if t.SessionData.Login == "" {
return trace.BadParameter("login: missing login")
}
if t.SessionData.ServerID == "" {
return trace.BadParameter("server: missing server")
}
if err := t.Term.CheckAndSetDefaults(); err != nil {
return trace.Wrap(err)
}
if t.UserAuthClient == nil {
return trace.BadParameter("UserAuthClient must be provided")
}
if t.LocalAccessPoint == nil {
return trace.BadParameter("localAccessPoint must be provided")
}
if t.SessionCtx == nil {
return trace.BadParameter("SessionCtx must be provided")
}
if t.Router == nil {
return trace.BadParameter("Router must be provided")
}
if t.TracerProvider == nil {
t.TracerProvider = tracing.DefaultProvider()
}
if t.Clock == nil {
t.Clock = clockwork.NewRealClock()
}
t.tracer = t.TracerProvider.Tracer("webterminal")
return nil
}
// sshBaseHandler is a base handler for web SSH connections.
type sshBaseHandler struct {
// logger holds the structured logger.
logger *slog.Logger
// ctx is a web session context for the currently logged-in user.
ctx *SessionContext
// userAuthClient is used to fetch nodes and sessions from the backend via the users' identity.
userAuthClient UserAuthClient
// proxyHostPort is the address of the server to connect to.
proxyHostPort string
// proxyPublicAddr is the public web proxy address.
proxyPublicAddr string
// keepAliveInterval is the interval for sending ping frames to a web client.
// This value is pulled from the cluster network config and
// guaranteed to be set to a nonzero value as it's enforced by the configuration.
keepAliveInterval time.Duration
// The server data for the active session.
sessionData session.Session
// router is used to dial the host
router *proxy.Router
// tracer creates spans
tracer oteltrace.Tracer
// localAccessPoint is the subset of the Proxy cache required to
// look up information from the local cluster. This should not
// be used for anything that requires RBAC on behalf of the user.
// Requests that should be made on behalf of the user should
// use [UserAuthClient].
localAccessPoint localAccessPoint
// interactiveCommand is a command to execute.
interactiveCommand []string
// resolver looks up the hostname for the server UUID.
resolver func(serverID string) (hostname string, err error)
// sshDialTimeout is the maximum time to wait for an SSH connection
// to be established before aborting.
sshDialTimeout time.Duration
// fipsBuild indicates if this is a Teleport FIPS build.
fipsBuild bool
}
// localAccessPoint is a subset of the cache used to look up
// various cluster details.
type localAccessPoint interface {
GetUser(ctx context.Context, username string, withSecrets bool) (types.User, error)
GetRole(ctx context.Context, name string) (types.Role, error)
}
// TerminalHandler connects together an SSH session with a web-based
// terminal via a web socket.
type TerminalHandler struct {
sshBaseHandler
// displayLogin is the login name to display in the UI.
displayLogin string
closeOnce sync.Once
// term is the initial PTY size.
term session.TerminalParams
// proxySigner is used to sign PROXY header and securely propagate client IP information
proxySigner multiplexer.PROXYHeaderSigner
// participantMode is the mode that determines what you can do when you join an active session.
participantMode types.SessionParticipantMode
// stream manages sending and receiving [Envelope] to the UI
// for the duration of the session
stream *terminal.Stream
// tracker is the session tracker of the session being joined. May be nil
// if the user is not joining a session.
tracker types.SessionTracker
// presenceChecker to use for presence checking
presenceChecker PresenceChecker
// closedByClient indicates if the websocket connection was closed by the
// user (closing the browser tab, exiting the session, etc).
closedByClient atomic.Bool
// clock used to interact with time.
clock clockwork.Clock
// websocketConn is the active websocket connection
websocketConn *websocket.Conn
}
// ServeHTTP builds a connection to the remote node and then pumps back two types of
// events: raw input/output events for what's happening on the terminal itself
// and audit log events relevant to this session.
func (t *TerminalHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// This allows closing of the websocket if the user logs out before exiting
// the session.
t.ctx.AddClosers(t)
defer t.ctx.RemoveCloser(t)
ws := t.websocketConn
err := ws.SetReadDeadline(deadlineForInterval(t.keepAliveInterval))
if err != nil {
t.logger.ErrorContext(r.Context(), "Error setting websocket readline", "error", err)
return
}
t.handler(ws, r)
}
func (t *TerminalHandler) writeSessionData(ctx context.Context) error {
envelope := &terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketSessionMetadata,
}
sessionDataTemp := t.sessionData
// If the displayLogin is set then use it in the session metadata instead of the
// login name used in the SSH connection. This is specifically for the use case
// when joining a session to avoid displaying "-teleport-internal-join" as the username.
if t.displayLogin != "" {
sessionDataTemp.Login = t.displayLogin
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp})
if err != nil {
t.sendError(ctx, "unable to marshal session response", err, t.stream)
return trace.Wrap(err)
}
envelope.Payload = string(sessionMetadataResponse)
} else {
// The Proxy cache is used to retrieve the server and resolve the hostname here instead
// of the user auth client to avoid a round trip to the Auth server. This would normally
// not be ok since this bypasses user RBAC, however, since at this point we have already
// established a connection to the target host via the user identity, the user MUST have
// access to the target host.
hostname, err := t.resolver(sessionDataTemp.ServerID)
if err != nil {
return trace.Wrap(err)
}
sessionDataTemp.ServerHostname = hostname
sessionMetadataResponse, err := json.Marshal(siteSessionGenerateResponse{Session: sessionDataTemp})
if err != nil {
t.sendError(ctx, "unable to marshal session response", err, t.stream)
return trace.Wrap(err)
}
envelope.Payload = string(sessionMetadataResponse)
}
envelopeBytes, err := proto.Marshal(envelope)
if err != nil {
t.sendError(ctx, "unable to marshal session data event for web client", err, t.stream)
return trace.Wrap(err)
}
if err := t.stream.WriteMessage(websocket.BinaryMessage, envelopeBytes); err != nil {
t.sendError(ctx, "unable to write message to socket", err, t.stream)
return trace.Wrap(err)
}
return nil
}
// Close the websocket stream.
func (t *TerminalHandler) Close() error {
var err error
t.closeOnce.Do(func() {
if t.stream == nil {
return
}
err = trace.Wrap(t.stream.Close())
})
return trace.Wrap(err)
}
// handler is the main websocket loop. It creates a Teleport client and then
// pumps raw events and audit events back to the client until the SSH session
// is complete.
func (t *TerminalHandler) handler(ws *websocket.Conn, r *http.Request) {
defer ws.Close()
// Update the read deadline upon receiving a pong message.
ws.SetPongHandler(func(_ string) error {
return trace.Wrap(ws.SetReadDeadline(deadlineForInterval(t.keepAliveInterval)))
})
// Create a context for signaling when the terminal session is over and
// link it first with the trace context from the request context
tctx := oteltrace.ContextWithRemoteSpanContext(context.Background(), oteltrace.SpanContextFromContext(r.Context()))
ctx, cancel := context.WithCancel(tctx)
defer cancel()
t.stream = terminal.NewStream(ctx, terminal.StreamConfig{WS: ws, Logger: t.logger})
// Create a Teleport client, if not able to, show the reason to the user in
// the terminal.
tc, err := t.makeClient(ctx, t.stream, ws.RemoteAddr().String())
if err != nil {
t.logger.InfoContext(ctx, "Failed creating a client for session", "error", err)
t.stream.WriteError(ctx, err.Error())
return
}
t.logger.DebugContext(ctx, "Creating websocket stream")
defaultCloseHandler := ws.CloseHandler()
ws.SetCloseHandler(func(code int, text string) error {
t.closedByClient.Store(true)
t.logger.DebugContext(ctx, "web socket was closed by client - terminating session")
// Call the default close handler if one was set.
if defaultCloseHandler != nil {
err := defaultCloseHandler(code, text)
return trace.NewAggregate(err, t.Close())
}
return trace.Wrap(t.Close())
})
// Start sending ping frames through websocket to client.
go startWSPingLoop(ctx, ws, t.keepAliveInterval, t.logger, t.Close)
// Pump raw terminal in/out and audit events into the websocket.
go t.streamEvents(ctx, tc)
// Block until the terminal session is complete.
t.streamTerminal(ctx, tc)
t.logger.DebugContext(ctx, "Closing websocket stream")
}
type stderrWriter struct {
stream *terminal.Stream
}
func (s stderrWriter) Write(b []byte) (int, error) {
s.stream.WriteError(context.Background(), string(b))
return len(b), nil
}
// makeClient builds a *client.TeleportClient for the connection.
func (t *TerminalHandler) makeClient(ctx context.Context, stream *terminal.Stream, clientAddr string) (*client.TeleportClient, error) {
ctx, span := tracing.DefaultProvider().Tracer("terminal").Start(ctx, "terminal/makeClient")
defer span.End()
clientConfig, err := makeTeleportClientConfig(ctx, t.ctx)
if err != nil {
return nil, trace.Wrap(err)
}
clientConfig.HostLogin = t.sessionData.Login
clientConfig.ForwardAgent = client.ForwardAgentLocal
clientConfig.Stdout = stream
clientConfig.Stderr = stderrWriter{stream: stream}
clientConfig.Stdin = stream
clientConfig.SiteName = t.sessionData.ClusterName
if err := clientConfig.ParseProxyHost(t.proxyHostPort); err != nil {
return nil, trace.BadParameter("failed to parse proxy address: %v", err)
}
clientConfig.Host = t.sessionData.ServerHostname
clientConfig.HostPort = t.sessionData.ServerHostPort
clientConfig.SessionID = t.sessionData.ID.String()
clientConfig.ClientAddr = clientAddr
clientConfig.Tracer = t.tracer
clientConfig.SSHDialTimeout = t.sshDialTimeout
if len(t.interactiveCommand) > 0 {
clientConfig.InteractiveCommand = true
}
tc, err := client.NewClient(clientConfig)
if err != nil {
return nil, trace.BadParameter("failed to create client: %v", err)
}
// Save the *ssh.Session after the shell has been created. The session is
// used to update all other parties window size to that of the web client and
// to allow future window changes.
tc.OnShellCreated = func(s *tracessh.Session, c *tracessh.Client, _ io.ReadWriteCloser) (bool, error) {
if err := t.stream.SessionCreated(s); err != nil {
t.logger.DebugContext(ctx, "terminating established ssh connection to host",
"error", err,
)
return false, trace.Wrap(s.Close())
}
// The web session was closed by the client while the ssh connection was being established.
// Attempt to close the SSH session instead of proceeding with the window change request.
if t.closedByClient.Load() {
t.logger.DebugContext(ctx, "websocket was closed by client, terminating established ssh connection to host")
return false, trace.Wrap(s.Close())
}
if err := s.WindowChange(ctx, t.term.H, t.term.W); err != nil {
t.logger.ErrorContext(ctx, "failed to send window change request", "error", err)
}
return false, nil
}
return tc, nil
}
// issueSessionMFACerts performs the mfa ceremony to retrieve new certs that can be
// used to access nodes which require per-session mfa. The ceremony is performed directly
// to make use of the userAuthClient already established for the session instead of leveraging
// the TeleportClient which would require dialing the auth server a second time.
func (t *sshBaseHandler) issueSessionMFACerts(ctx context.Context, tc *client.TeleportClient, wsStream *terminal.WSStream) ([]ssh.Signer, error) {
ctx, span := t.tracer.Start(ctx, "terminal/issueSessionMFACerts")
defer span.End()
t.logger.DebugContext(ctx, "Attempting to issue a single-use user certificate with an MFA check")
// Prepare MFA check request.
mfaRequiredReq := &authproto.IsMFARequiredRequest{
Target: &authproto.IsMFARequiredRequest_Node{
Node: &authproto.NodeLogin{
Node: t.sessionData.ServerID,
Login: tc.HostLogin,
},
},
}
// Prepare UserCertsRequest.
pk, err := keys.ParsePrivateKey(t.ctx.cfg.Session.GetSSHPriv())
if err != nil {
return nil, trace.Wrap(err)
}
sshCert, err := sshutils.ParseCertificate(t.ctx.cfg.Session.GetPub())
if err != nil {
return nil, trace.Wrap(err)
}
expires := time.Unix(int64(sshCert.ValidBefore), 0)
certsReq := &authproto.UserCertsRequest{
SSHPublicKey: pk.MarshalSSHPublicKey(),
Username: sshCert.KeyId, // SSH cert KeyId is set to teleport username.
Expires: expires,
RouteToCluster: t.sessionData.ClusterName,
NodeName: t.sessionData.ServerID,
Usage: authproto.UserCertsRequest_SSH,
Format: tc.CertificateFormat,
SSHLogin: tc.HostLogin,
}
result, err := client.PerformSessionMFACeremony(ctx, client.PerformSessionMFACeremonyParams{
CurrentAuthClient: t.userAuthClient,
RootAuthClient: t.ctx.cfg.RootClient,
MFACeremony: newMFACeremony(wsStream, t.ctx.cfg.RootClient.CreateAuthenticateChallenge, t.proxyPublicAddr),
MFAAgainstRoot: t.ctx.cfg.RootClusterName == tc.SiteName,
MFARequiredReq: mfaRequiredReq,
CertsReq: certsReq,
})
if err != nil {
return nil, trace.Wrap(err)
}
sshCert, err = sshutils.ParseCertificate(result.NewCerts.SSH)
if err != nil {
return nil, trace.Wrap(err)
}
signer, err := sshutils.SSHSigner(sshCert, pk)
if err != nil {
return nil, trace.Wrap(err)
}
return []ssh.Signer{signer}, nil
}
func newMFACeremony(stream *terminal.WSStream, createAuthenticateChallenge mfa.CreateAuthenticateChallengeFunc, proxyAddr string) *mfa.Ceremony {
// channelID is used by the front end to differentiate between separate ongoing SSO challenges.
channelID := uuid.NewString()
return &mfa.Ceremony{
CreateAuthenticateChallenge: createAuthenticateChallenge,
MFACeremonyConstructor: func(context.Context) (mfa.CallbackCeremony, error) {
return newMFACallbackCeremony(channelID, proxyAddr)
},
PromptConstructor: newMFAPromptConstructor(stream, channelID),
}
}
func newMFACeremonyPerformer(stream *terminal.Stream, mfaClient mfav2.MFAServiceClient, proxyAddr string, targetCluster string) (clientssh.MFACeremonyPerformer, error) {
// channelID is used by the front end to differentiate between separate ongoing SSO challenges.
channelID := uuid.NewString()
config := mfa.SessionBoundCeremonyConfig{
CreateSessionChallenge: mfaClient.CreateSessionChallenge,
ValidateSessionChallenge: mfaClient.ValidateSessionChallenge,
PromptConstructor: newMFAPromptConstructor(stream.WSStream, channelID),
CallbackCeremonyConstructor: func(context.Context) (mfa.CallbackCeremony, error) {
return newMFACallbackCeremony(channelID, proxyAddr)
},
TargetCluster: targetCluster,
}
ceremony, err := mfa.NewSessionBoundCeremony(config)
if err != nil {
return nil, trace.Wrap(err)
}
return func(ctx context.Context, sessionID []byte) (string, error) {
name, err := ceremony.Run(
ctx,
mfav2.SessionIdentifyingPayload_builder{
SshSessionId: sessionID,
}.Build(),
)
if err != nil {
return "", trace.Wrap(err)
}
return name, nil
}, nil
}
func newMFAPromptConstructor(stream *terminal.WSStream, channelID string) mfa.PromptConstructor {
return func(...mfa.PromptOpt) mfa.Prompt {
return mfa.PromptFunc(func(ctx context.Context, chal *authproto.MFAAuthenticateChallenge) (*authproto.MFAAuthenticateResponse, error) {
// Convert from proto to JSON types.
var challenge client.MFAAuthenticateChallenge
if chal.WebauthnChallenge != nil {
challenge.WebauthnChallenge = wantypes.CredentialAssertionFromProto(chal.WebauthnChallenge)
}
if chal.SSOChallenge != nil {
challenge.SSOChallenge = client.SSOChallengeFromProto(chal.SSOChallenge)
challenge.SSOChallenge.ChannelID = channelID
}
if chal.WebauthnChallenge == nil && chal.SSOChallenge == nil {
return nil, trace.Wrap(authclient.ErrNoMFADevices)
}
var codec protobufMFACodec
if err := stream.WriteChallenge(&challenge, codec); err != nil {
return nil, trace.Wrap(err)
}
resp, err := stream.ReadChallengeResponse(codec)
return resp, trace.Wrap(err)
})
}
}
func newMFACallbackCeremony(channelID string, proxyAddr string) (mfa.CallbackCeremony, error) {
u, err := url.Parse(sso.WebMFARedirect)
if err != nil {
return nil, trace.Wrap(err)
}
u.RawQuery = url.Values{"channel_id": {channelID}}.Encode()
return &sso.MFACeremony{
ClientCallbackURL: u.String(),
ProxyAddress: proxyAddr,
}, nil
}
type connectWithMFAFn = func(ctx context.Context, scopePin *scopesv1.Pin, stream *terminal.Stream, tc *client.TeleportClient, accessChecker services.AccessChecker, getAgent sshagent.ClientGetter, signer agentless.SignerCreator) (*client.NodeClient, error)
// connectToHost establishes a connection to the target host. It first attempts to connect with the existing
// certificates, which can succeed if per-session MFA is not required or if the node supports in-band MFA. If that
// fails, it falls back to the legacy per-session MFA certificate flow. If both attempts fail, it returns the error
// that is most likely to be helpful to the user.
func (t *sshBaseHandler) connectToHost(ctx context.Context, stream *terminal.Stream, tc *client.TeleportClient, connectToNodeWithMFA connectWithMFAFn) (*client.NodeClient, error) {
ctx, span := t.tracer.Start(ctx, "terminal/connectToHost")
defer span.End()
accessChecker, err := t.ctx.GetUserAccessChecker()
if err != nil {
return nil, trace.Wrap(err)
}
getAgent := sshagent.NewStaticClientGetter(tc.LocalAgent())
cert, err := t.ctx.GetSSHCertificate()
if err != nil {
return nil, trace.Wrap(err)
}
ident, err := sshca.DecodeIdentity(cert)
if err != nil {
return nil, trace.Wrap(err)
}
certGen, err := t.router.GetSiteClient(ctx, tc.SiteName)
if err != nil {
return nil, trace.Wrap(err)
}
signer := agentless.SignerFromSSHIdentity(ident, t.localAccessPoint, certGen, tc.SiteName, tc.Username)
// Try to connect directly with existing certs. This can succeed if MFA is not required or if the node supports
// in-band MFA and the ceremony completes successfully.
clt, directErr := t.connectToNode(ctx, ident.ScopePin, stream, tc, accessChecker, getAgent, signer)
if directErr == nil {
return clt, nil
}
// Fall back to attempting to connect with certs issued from the session MFA ceremony. This can succeed if MFA is
// required and the user has enrolled MFA devices, but will fail if the user does not have any enrolled MFA devices
// or if there are any issues during the MFA ceremony.
//
// TODO(cthach): DELETE IN v20.0 when the legacy per-session MFA with certifcates flow is removed.
clt, mfaErr := connectToNodeWithMFA(ctx, ident.ScopePin, stream, tc, accessChecker, getAgent, signer)
if mfaErr == nil {
return clt, nil
}
switch {
// Any direct connection errors other than access denied, which should be returned
// if MFA is required, take precedent over MFA errors due to users not having any
// enrolled devices.
case !trace.IsAccessDenied(directErr) && errors.Is(mfaErr, authclient.ErrNoMFADevices):
return nil, trace.Wrap(directErr)
case !errors.Is(mfaErr, io.EOF) && // Ignore any errors from MFA due to locks being enforced, the direct error will be friendlier
!errors.As(mfaErr, new(*client.MFARequiredUnknownError)) && // Ignore any failures that occurred before determining if MFA was required
!errors.Is(mfaErr, services.ErrSessionMFANotRequired): // Ignore any errors caused by attempting the MFA ceremony when MFA will not grant access
return nil, trace.Wrap(mfaErr)
default:
return nil, trace.Wrap(directErr)
}
}
// streamTerminal opens an SSH connection to the remote host and streams
// events back to the web client.
func (t *TerminalHandler) streamTerminal(ctx context.Context, tc *client.TeleportClient) {
ctx, span := t.tracer.Start(ctx, "terminal/streamTerminal")
defer span.End()
nc, err := t.connectToHost(ctx, t.stream, tc, t.connectToNodeWithMFA)
if err != nil {
t.logger.WarnContext(ctx, "Unable to stream terminal - failure connecting to host", "error", err)
t.stream.WriteError(ctx, err.Error())
return
}
defer nc.Close()
// If the session was terminated by client while the connection to the host
// was being established, then return early before creating the shell. Any terminations
// by the client from here on out should either get caught in the OnShellCreated callback
// set on the [tc] or in [TerminalHandler.Close].
if t.closedByClient.Load() {
t.logger.DebugContext(ctx, "websocket was closed by client, aborting establishing ssh connection to host")
return
}
var beforeStart func(io.Writer)
if t.participantMode == types.SessionModeratorMode {
beforeStart = func(out io.Writer) {
nc.OnMFA = func() {
baseCeremony := newMFACeremony(t.stream.WSStream, nil, t.proxyPublicAddr)
if err := t.presenceChecker(ctx, out, t.userAuthClient, t.sessionData.ID.String(), baseCeremony); err != nil {
t.logger.WarnContext(ctx, "Unable to stream terminal - failure performing presence checks", "error", err)
return
}
}
}
}
monitorCtx, monitorCancel := context.WithCancel(ctx)
defer monitorCancel()
sshPinger, err := latency.NewSSHPinger(nc.Client)
if err != nil {
t.logger.WarnContext(monitorCtx, "failure monitoring session latency", "error", err)
} else {
go monitorLatency(monitorCtx, t.clock, t.stream.WSStream, sshPinger,
latency.ReporterFunc(
func(ctx context.Context, statistics latency.Statistics) error {
return trace.Wrap(
t.stream.WSStream.WriteLatency(terminal.SSHSessionLatencyStats{
WebSocket: statistics.Client,
SSH: statistics.Server,
}),
)
},
),
)
}
sessionDataSent := make(chan struct{})
// If we are joining a session, send the session data right away, we
// know the session ID
if t.tracker != nil {
if err := t.writeSessionData(ctx); err != nil {
t.logger.WarnContext(ctx, "Failure sending session data", "error", err)
}
close(sessionDataSent)
} else {
// We are creating a new session and the server will generate a
// new session ID, send the session data once the session is
// created and the server sends us the session ID it is using
writeSessionCtx, writeSessionCancel := context.WithCancel(ctx)
defer writeSessionCancel()
// only handle the first session ID request
var receiveSessionIDOnce sync.Once
receivedSessionID := make(chan struct{})
nc.Client.HandleSessionRequest(ctx, teleport.CurrentSessionIDRequest, func(ctx context.Context, req *ssh.Request) {
receiveSessionIDOnce.Do(func() {
sid, err := session.ParseID(string(req.Payload))
if err != nil {
t.logger.WarnContext(ctx, "Unable to parse session ID", "error", err)
return
}
t.sessionData.ID = *sid
close(receivedSessionID)
})
})
// wait in a new goroutine because the server won't set a
// session ID until we start the session.
go func() {
defer close(sessionDataSent)
ctx, cancel := context.WithTimeout(writeSessionCtx, 10*time.Second)
defer cancel()
select {
case <-receivedSessionID:
if err := t.writeSessionData(ctx); err != nil {
t.logger.WarnContext(ctx, "Failure sending session data", "error", err)
}
case <-ctx.Done():
t.logger.WarnContext(ctx, "Failed to receive session data")
}
}()
}
var joinSessionID string
if t.tracker != nil {
joinSessionID = t.tracker.GetSessionID()
}
// Establish SSH connection to the server. This function will block until
// either an error occurs or it completes successfully.
if err = nc.RunInteractiveShell(ctx, joinSessionID, t.participantMode, beforeStart); err != nil {
if !t.closedByClient.Load() {
t.stream.WriteError(ctx, err.Error())
}
return
}
if t.closedByClient.Load() {
return
}
// Wait for the session data to be sent before closing the session
<-sessionDataSent
// Send close envelope to web terminal upon exit without an error.
if err := t.stream.SendCloseMessage(t.sessionData.ServerID); err != nil {
t.logger.ErrorContext(ctx, "Unable to send close event to web client", "error", err)
}
if err := t.stream.Close(); err != nil && !errors.Is(err, io.EOF) {
t.logger.ErrorContext(ctx, "Unable to close client web socket", "error", err)
return
}
t.logger.DebugContext(ctx, "Sent close event to web client")
}
// generateClientConfig creates an [apissh.ClientConfig] for the Teleport client connection.
func (t *sshBaseHandler) generateClientConfig(ctx context.Context, stream *terminal.Stream, tc *client.TeleportClient) (apissh.ClientConfig, error) {
performer, err := newMFACeremonyPerformer(
stream,
t.userAuthClient.MFAServiceClientV2(),
t.proxyPublicAddr,
tc.SiteName,
)
if err != nil {
return apissh.ClientConfig{}, trace.Wrap(err)
}
authCallback := clientssh.AuthCallback(
ctx,
clientssh.AuthCallbackConfig{
MFAPerformer: performer,
},
)
return apissh.ClientConfig{
User: tc.HostLogin,
PublicKeyAuth: tc.PublicKeyAuthConfig,
AuthCallback: authCallback,
HostKeyCallback: tc.HostKeyCallback,
Timeout: t.sshDialTimeout,
}, nil
}
// connectToNode attempts to connect to the host with the already
// provisioned certs for the user.
func (t *sshBaseHandler) connectToNode(ctx context.Context, scopePin *scopesv1.Pin, stream *terminal.Stream, tc *client.TeleportClient, accessChecker services.AccessChecker, getAgent sshagent.ClientGetter, signer agentless.SignerCreator) (*client.NodeClient, error) {
conn, err := t.router.DialHost(ctx, scopePin, stream.RemoteAddr(), stream.LocalAddr(), t.sessionData.ServerID, strconv.Itoa(t.sessionData.ServerHostPort), tc.SiteName, accessChecker.CheckAccessToRemoteCluster, getAgent, signer)
if err != nil {
t.logger.WarnContext(ctx, "Unable to stream terminal - failed to dial host", "error", err)
if errors.Is(err, teleport.ErrNodeIsAmbiguous) {
const message = "error: ambiguous host could match multiple nodes\n\nHint: try addressing the node by unique id (ex: user@node-id)\n"
return nil, trace.NotFound("%s", message)
}
return nil, trace.Wrap(err)
}
sshConfig, err := t.generateClientConfig(ctx, stream, tc)
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := client.NewNodeClient(ctx, sshConfig, conn,
net.JoinHostPort(t.sessionData.ServerID, strconv.Itoa(t.sessionData.ServerHostPort)),
t.sessionData.ServerHostname,
tc, t.fipsBuild)
if err != nil {
// The close error is ignored instead of using [trace.NewAggregate] because
// aggregate errors do not allow error inspection with things like [trace.IsAccessDenied].
_ = conn.Close()
// Since connection attempts are made via UUID and not hostname, any access denied errors
// will not contain the resolved host address. To provide an easier troubleshooting experience
// for users, attempt to resolve the hostname of the server and augment the error message with it.
if trace.IsAccessDenied(err) {
if resp, err := t.userAuthClient.ListUnifiedResources(ctx, &authproto.ListUnifiedResourcesRequest{
SortBy: types.SortBy{Field: types.ResourceKind},
Kinds: []string{types.KindNode},
Limit: 1,
PredicateExpression: fmt.Sprintf(`resource.metadata.name == %q`, t.sessionData.ServerID),
}); err == nil && len(resp.Resources) > 0 {
return nil, trace.AccessDenied("access denied to %q connecting to %v", sshConfig.User, resp.Resources[0].GetNode().GetHostname())
}
}
return nil, trace.Wrap(err)
}
clt.ProxyPublicAddr = t.proxyPublicAddr
return clt, nil
}
// connectToNodeWithMFA attempts to perform the mfa ceremony and then dial the
// host with the retrieved single use certs.
func (t *TerminalHandler) connectToNodeWithMFA(ctx context.Context, scopePin *scopesv1.Pin, stream *terminal.Stream, tc *client.TeleportClient, accessChecker services.AccessChecker, getAgent sshagent.ClientGetter, signer agentless.SignerCreator) (*client.NodeClient, error) {
// perform mfa ceremony and retrieve new certs
signers, err := t.issueSessionMFACerts(ctx, tc, stream.WSStream)
if err != nil {
return nil, trace.Wrap(err)
}
return t.connectToNodeWithMFABase(ctx, scopePin, stream, tc, accessChecker, getAgent, signer, signers)
}
// connectToNodeWithMFABase attempts to dial the host with the provided auth
// methods.
func (t *sshBaseHandler) connectToNodeWithMFABase(
ctx context.Context,
scopePin *scopesv1.Pin,
stream *terminal.Stream,
tc *client.TeleportClient,
accessChecker services.AccessChecker,
getAgent sshagent.ClientGetter,
agentlessSigner agentless.SignerCreator,
signers []ssh.Signer,
) (*client.NodeClient, error) {
sshConfig := apissh.ClientConfig{
User: tc.HostLogin,
PublicKeyAuth: apissh.PublicKeyAuthConfig{
Signers: func() ([]ssh.Signer, error) {
return signers, nil
},
},
HostKeyCallback: tc.HostKeyCallback,
Timeout: t.sshDialTimeout,
}
// connect to the node again with the new certs
conn, err := t.router.DialHost(ctx, scopePin, stream.RemoteAddr(), stream.LocalAddr(), t.sessionData.ServerID, strconv.Itoa(t.sessionData.ServerHostPort), tc.SiteName, accessChecker.CheckAccessToRemoteCluster, getAgent, agentlessSigner)
if err != nil {
return nil, trace.Wrap(err)
}
nc, err := client.NewNodeClient(ctx, sshConfig, conn,
net.JoinHostPort(t.sessionData.ServerID, strconv.Itoa(t.sessionData.ServerHostPort)),
t.sessionData.ServerHostname,
tc, t.fipsBuild)
if err != nil {
return nil, trace.NewAggregate(err, conn.Close())
}
nc.ProxyPublicAddr = t.proxyPublicAddr
return nc, nil
}
// sendError sends an error message to the client using the provided websocket.
func (t *sshBaseHandler) sendError(ctx context.Context, errMsg string, err error, ws terminal.WSConn) {
envelope := &terminal.Envelope{
Version: defaults.WebsocketVersion,
Type: defaults.WebsocketError,
Payload: fmt.Sprintf("%s: %s", errMsg, err.Error()),
}
envelopeBytes, err := proto.Marshal(envelope)
if err != nil {
t.logger.ErrorContext(ctx, "failed to marshal error message", "error", err)
}
if err := ws.WriteMessage(websocket.BinaryMessage, envelopeBytes); err != nil {
t.logger.ErrorContext(ctx, "failed to send error message", "error", err)
}
}
// streamEvents receives events over the SSH connection and forwards them to
// the web client.
func (t *TerminalHandler) streamEvents(ctx context.Context, tc *client.TeleportClient) {
for {
select {
// Send push events that come over the events channel to the web client.
case event := <-tc.EventsChannel():
logger := t.logger.With("event", event.GetType())
data, err := json.Marshal(event)
if err != nil {
logger.ErrorContext(ctx, "Unable to marshal audit event", "error", err)
continue
}
logger.DebugContext(ctx, "Sending audit event to web client")
if err := t.stream.WriteAuditEvent(data); err != nil {
if errors.Is(err, websocket.ErrCloseSent) {
logger.DebugContext(ctx, "Websocket was closed, no longer streaming events", "error", err)
return
}
logger.ErrorContext(ctx, "Unable to send audit event to web client", "error", err)
continue
}
// Once the terminal stream is over (and the close envelope has been sent),
// close stop streaming envelopes.
case <-ctx.Done():
return
}
}
}
// the defaultPort of 0 indicates that the port is
// unknown or was not provided and should be guessed
const defaultPort = 0
// resolveServerHostPort parses server name and attempts to resolve hostname
// and port.
func resolveServerHostPort(servername string, existingServers []types.Server) (string, int, error) {
if servername == "" {
return "", defaultPort, trace.BadParameter("empty server name")
}
// Check if servername is UUID.
for _, node := range existingServers {
if node.GetName() == servername {
return node.GetHostname(), defaultPort, nil
}
}
host, port, err := serverHostPort(servername)
return host, port, trace.Wrap(err)
}
// serverHostPort returns the host and port for [servername]
func serverHostPort(servername string) (string, int, error) {
if !strings.Contains(servername, ":") {
return servername, defaultPort, nil
}
// Check for explicitly specified port.
host, portString, err := utils.SplitHostPort(servername)
if err != nil {
return "", defaultPort, trace.Wrap(err)
}
port, err := strconv.Atoi(portString)
if err != nil {
return "", defaultPort, trace.BadParameter("invalid port: %v", err)
}
return host, port, nil
}
// deadlineForInterval returns a suitable network read deadline for a given ping interval.
// We chose to take the current time plus twice the interval to allow the timeframe of one interval
// to wait for a returned pong message.
func deadlineForInterval(interval time.Duration) time.Time {
return time.Now().Add(interval * 2)
}
/*
* 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 (
"bytes"
"context"
"encoding/binary"
"net/http"
"time"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/types/events"
"github.com/gravitational/teleport/lib/player"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/session"
"github.com/gravitational/teleport/lib/utils"
logutils "github.com/gravitational/teleport/lib/utils/log"
)
const (
messageTypePTY = byte(1)
messageTypeError = byte(2)
messageTypePlayPause = byte(3)
messageTypeSeek = byte(4)
messageTypeResize = byte(5)
)
const (
severityError = byte(1)
)
const (
actionPlay = byte(0)
actionPause = byte(1)
)
func (h *Handler) sessionLengthHandle(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
) (any, error) {
sID := p.ByName("sid")
if sID == "" {
return nil, trace.BadParameter("missing session ID in request URL")
}
ctx, cancel := context.WithCancel(r.Context())
defer cancel()
clt, err := sctx.GetUserClient(ctx, cluster)
if err != nil {
return nil, trace.Wrap(err)
}
type response struct {
Duration int64 `json:"durationMs"`
RecordingType string `json:"recordingType"`
}
evts, errs := clt.StreamSessionEvents(ctx, session.ID(sID), 0)
for {
select {
case err := <-errs:
return nil, trace.Wrap(err)
case evt, ok := <-evts:
if !ok {
return nil, trace.NotFound("could not find end event for session %v", sID)
}
switch evt := evt.(type) {
case *events.SessionEnd:
return response{evt.EndTime.Sub(evt.StartTime).Milliseconds(), "ssh"}, nil
case *events.WindowsDesktopSessionEnd:
return response{evt.EndTime.Sub(evt.StartTime).Milliseconds(), "desktop"}, nil
case *events.DatabaseSessionEnd:
return response{evt.EndTime.Sub(evt.StartTime).Milliseconds(), "database"}, nil
}
}
}
}
func (h *Handler) ttyPlaybackHandle(
w http.ResponseWriter,
r *http.Request,
p httprouter.Params,
sctx *SessionContext,
cluster reversetunnelclient.Cluster,
) (any, error) {
sID := p.ByName("sid")
if sID == "" {
return nil, trace.BadParameter("missing session ID in request URL")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
h.logger.DebugContext(r.Context(), "upgrading to websocket")
upgrader := websocket.Upgrader{
ReadBufferSize: 1024,
WriteBufferSize: 1024,
}
ws, err := upgrader.Upgrade(w, r, nil)
if err != nil {
h.logger.WarnContext(r.Context(), "failed upgrade", "error", err)
// if Upgrade fails, it automatically replies with an HTTP error
// (this means we don't need to return an error here)
return nil, nil
}
player, err := player.New(&player.Config{
Clock: h.clock,
Log: h.logger,
SessionID: session.ID(sID),
Streamer: clt,
Context: r.Context(),
})
if err != nil {
h.logger.WarnContext(r.Context(), "player error", "error", err)
writeError(ws, err)
return nil, nil
}
ctx, cancel := context.WithCancel(r.Context())
defer cancel()
go func() {
defer cancel()
for {
typ, b, err := ws.ReadMessage()
if err != nil {
if !utils.IsOKNetworkError(err) {
h.logger.WarnContext(ctx, "websocket read error", "error", err)
}
return
}
if typ != websocket.BinaryMessage {
h.logger.DebugContext(ctx, "skipping unknown websocket message type", "message_type", logutils.TypeAttr(typ))
continue
}
if err := handlePlaybackAction(b, player); err != nil {
h.logger.WarnContext(ctx, "skipping bad action", "error", err)
continue
}
}
}()
go func() {
defer cancel()
defer func() {
h.logger.DebugContext(ctx, "closing websocket")
if err := ws.WriteMessage(websocket.CloseMessage, nil); err != nil {
h.logger.DebugContext(r.Context(), "error sending close message", "error", err)
}
if err := ws.Close(); err != nil {
h.logger.DebugContext(ctx, "error closing websocket", "error", err)
}
}()
player.Play()
defer player.Close()
headerBuf := make([]byte, 11)
headerBuf[0] = messageTypePTY
writePTY := func(b []byte, delay uint64) error {
writer, err := ws.NextWriter(websocket.BinaryMessage)
if err != nil {
return trace.Wrap(err, "getting websocket writer")
}
msgLen := uint16(len(b) + 8)
binary.BigEndian.PutUint16(headerBuf[1:], msgLen)
binary.BigEndian.PutUint64(headerBuf[3:], delay)
if _, err := writer.Write(headerBuf); err != nil {
return trace.Wrap(err, "writing message header")
}
// TODO(zmb3): consider optimizing this by bufering for very large sessions
// (wait up to N ms to batch events into a single websocket write).
if _, err := writer.Write(b); err != nil {
return trace.Wrap(err, "writing PTY data")
}
if err := writer.Close(); err != nil {
return trace.Wrap(err, "closing websocket writer")
}
return nil
}
writeSize := func(size string) error {
ts, err := session.UnmarshalTerminalParams(size)
if err != nil {
h.logger.DebugContext(ctx, "Ignoring invalid terminal size", "terminal_size", size)
return nil // don't abort playback due to a bad event
}
msg := make([]byte, 7)
msg[0] = messageTypeResize
binary.BigEndian.PutUint16(msg[1:], 4)
binary.BigEndian.PutUint16(msg[3:], uint16(ts.W))
binary.BigEndian.PutUint16(msg[5:], uint16(ts.H))
return trace.Wrap(ws.WriteMessage(websocket.BinaryMessage, msg))
}
for {
select {
case <-ctx.Done():
return
case evt, ok := <-player.C():
if !ok {
// send any playback errors to the browser
if err := writeError(ws, player.Err()); err != nil {
h.logger.WarnContext(ctx, "failed to send error message to browser", "error", err)
}
return
}
switch evt := evt.(type) {
case *events.SessionStart:
if err := writeSize(evt.TerminalSize); err != nil {
h.logger.DebugContext(ctx, "Failed to write resize", "error", err)
return
}
case *events.SessionPrint:
if err := writePTY(evt.Data, uint64(evt.DelayMilliseconds)); err != nil {
h.logger.DebugContext(ctx, "Failed to send PTY data", "error", err)
return
}
case *events.SessionEnd:
// send empty PTY data - this will ensure that any dead time
// at the end of the recording is processed and the allow
// the progress bar to go to 100%
if err := writePTY(nil, uint64(evt.EndTime.Sub(evt.StartTime)/time.Millisecond)); err != nil {
h.logger.DebugContext(ctx, "Failed to send session end data", "error", err)
return
}
case *events.Resize:
if err := writeSize(evt.TerminalSize); err != nil {
h.logger.DebugContext(ctx, "Failed to write resize", "error", err)
return
}
case *events.SessionLeave: // do nothing
default:
h.logger.DebugContext(ctx, "unexpected event type", "event_type", logutils.TypeAttr(evt))
}
}
}
}()
<-ctx.Done()
return nil, nil
}
func writeError(ws *websocket.Conn, err error) error {
if err == nil {
return nil
}
b := new(bytes.Buffer)
b.WriteByte(messageTypeError)
msg := trace.UserMessage(err)
l := 1 /* severity */ + 2 /* msg length */ + len(msg)
binary.Write(b, binary.BigEndian, uint16(l))
b.WriteByte(severityError)
binary.Write(b, binary.BigEndian, uint16(len(msg)))
b.WriteString(msg)
return trace.Wrap(ws.WriteMessage(websocket.BinaryMessage, b.Bytes()))
}
type play interface {
Play() error
Pause() error
SetPos(time.Duration) error
}
// handlePlaybackAction processes a playback message
// received from the browser
func handlePlaybackAction(b []byte, p play) error {
if len(b) < 3 {
return trace.BadParameter("invalid playback message")
}
msgType := b[0]
msgLen := binary.BigEndian.Uint16(b[1:])
if len(b) < int(msgLen)+3 {
return trace.BadParameter("invalid message length")
}
payload := b[3:]
payload = payload[:msgLen]
switch msgType {
case messageTypePlayPause:
if len(payload) != 1 {
return trace.BadParameter("invalid play/pause command")
}
switch action := payload[0]; action {
case actionPlay:
p.Play()
case actionPause:
p.Pause()
default:
return trace.BadParameter("invalid play/pause action %v", action)
}
case messageTypeSeek:
if len(payload) != 8 {
return trace.BadParameter("invalid seek message")
}
pos := binary.BigEndian.Uint64(payload)
p.SetPos(time.Duration(pos) * time.Millisecond)
}
return nil
}
/*
# Websocket Protocol for TTY Playback:
During playback, the Teleport proxy sends session data to the browser
and the browser sends playback commands (play/pause, seek, etc) to the
proxy.
Each message conforms to the following binary protocol.
## Message Header
The message header starts with a 1-byte identifier followed by a 2-byte
(big endian) integer containing the number of bytes following the header.
This length field does not include the 3-byte header.
## Messages
### 1 - PTY data
This message is used to send recorded PTY data to the browser.
- Message ID: 1
- 8-byte timestamp (milliseconds since session start)
- PTY data
### 2 - Error
This message is used to indicate that an error has occurred.
- Message ID: 2
- 1 byte severity (1=error)
- 2-byte error message length
- variable length error message (UTF-8 text)
### 3 - Play/Pause
This message is sent from the browser to the server to pause
or resume playback.
- Message ID: 3
- 1-byte code (0=play, 1=pause)
### 4 - Seek
This message is used to seek to a new position in the recording.
- Message ID: 4
- 8-byte timestamp (milliseconds since session start)
### 5 - Resize
This message is used to indicate that the terminal was resized.
- Message ID: 5
- 2-byte width
- 2-byte height
*/
/*
* Teleport
* Copyright (C) 2026 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 (
"github.com/gravitational/trace"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/services"
"github.com/gravitational/teleport/lib/utils/set"
webui "github.com/gravitational/teleport/lib/web/ui"
)
// UnifiedResourcePrincipals holds per-dimension principal sets for a unified
// resource. Only the fields relevant to the resource kind will be populated.
type UnifiedResourcePrincipals struct {
// Logins is populated for SSH nodes.
Logins *webui.PrincipalSet
// AWSRoleARNs is populated for AWS Console apps.
AWSRoleARNs *webui.PrincipalSet
}
// PrincipalsForUnifiedResourceOpts configures PrincipalsForUnifiedResource.
type PrincipalsForUnifiedResourceOpts struct {
// Resource is the enriched resource from the unified resource listing.
Resource *types.EnrichedResource
// CertPrincipals are the principals from the user's current certificate
// (used to filter SSH logins to those the cert can actually use).
CertPrincipals []string
// AccessChecker is the user's base AccessChecker.
AccessChecker services.AccessChecker
// IncludeRequestable indicates the response should distinguish between
// granted and requestable principals. When false, Granted == All.
IncludeRequestable bool
// UseSearchAsRoles indicates the request was made with search_as_roles,
// meaning enriched logins may include requestable principals.
UseSearchAsRoles bool
}
// PrincipalsForUnifiedResource computes the granted and requestable principals
// for a unified resource, based on the resource kind.
func PrincipalsForUnifiedResource(opts PrincipalsForUnifiedResourceOpts) (*UnifiedResourcePrincipals, error) {
result := &UnifiedResourcePrincipals{}
switch r := opts.Resource.ResourceWithLabels.(type) {
case types.Server:
logins, err := sshPrincipals(opts, r)
if err != nil {
return nil, trace.Wrap(err)
}
result.Logins = logins
case types.AppServer:
arns, err := appPrincipals(opts, r)
if err != nil {
return nil, trace.Wrap(err)
}
result.AWSRoleARNs = arns
}
return result, nil
}
// sshPrincipals computes login principals for an SSH node.
//
// When search_as_roles is active (UseSearchAsRoles or IncludeRequestable),
// enriched logins may contain requestable logins not in the user's certificate.
// These are returned as-is since they're for display or access-request
// purposes, not direct SSH connections.
//
// When IncludeRequestable is set, granted logins come from Auth's principal
// sets when present, else the base access checker, filtered to cert
// principals.
//
// In the default mode (neither flag set), all logins are filtered to cert
// principals so the connect menu only offers logins that will work.
func sshPrincipals(opts PrincipalsForUnifiedResourceOpts, server types.Server) (*webui.PrincipalSet, error) {
if opts.UseSearchAsRoles || opts.IncludeRequestable {
all := set.New(opts.Resource.Logins...)
ps := &webui.PrincipalSet{All: all}
if opts.IncludeRequestable {
granted, err := grantedLoginsForResource(opts, types.PrincipalTypeLogins, server)
if err != nil {
return nil, trace.Wrap(err)
}
ps.Granted = filterByIdentityPrincipals(opts.CertPrincipals, granted)
} else {
ps.Granted = all
}
return ps, nil
}
filtered := filterByIdentityPrincipals(opts.CertPrincipals, opts.Resource.Logins)
return &webui.PrincipalSet{All: filtered, Granted: filtered}, nil
}
// appPrincipals computes AWS role ARN principals for an app resource.
//
// AccessChecker's [GetAllowedLoginsForResource] is used for backward compatibility
// in case Auth does not support enriched resources.
func appPrincipals(opts PrincipalsForUnifiedResourceOpts, appServer types.AppServer) (*webui.PrincipalSet, error) {
// Get all visible ARNs (granted ∪ requestable).
all := opts.Resource.Logins
if len(all) == 0 {
var err error
all, err = opts.AccessChecker.GetAllowedLoginsForResource(appServer.GetApp())
if err != nil {
return nil, trace.Wrap(err)
}
}
allSet := set.New(all...)
ps := &webui.PrincipalSet{All: allSet}
if opts.IncludeRequestable {
granted, err := grantedLoginsForResource(opts, types.PrincipalTypeRoleARNs, appServer.GetApp())
if err != nil {
return nil, trace.Wrap(err)
}
ps.Granted = set.New(granted...)
} else {
ps.Granted = allSet
}
return ps, nil
}
// grantedLoginsForResource returns the logins usable without an access
// request, preferring Auth's precomputed granted set when present,
// else, computing locally using the base access checker.
//
// The local path covers a Proxy reading from an Auth that predates the
// principal sets, which happens while a cluster is part way through an
// upgrade.
//
// TODO(kiosion): DELETE IN 20.0.0
func grantedLoginsForResource(opts PrincipalsForUnifiedResourceOpts, kind string, resource services.AccessCheckable) ([]string, error) {
for _, ps := range opts.Resource.Principals {
if ps.PrincipalType == kind {
return ps.Granted, nil
}
}
granted, err := opts.AccessChecker.GetAllowedLoginsForResource(resource)
return granted, trace.Wrap(err)
}
// filterByIdentityPrincipals returns the intersection of allowedLogins with
// identityPrincipals as a set. This is equivalent to client.CalculateSSHLogins.
func filterByIdentityPrincipals(identityPrincipals, allowedLogins []string) set.Set[string] {
allowed := set.New(allowedLogins...)
result := set.NewWithCapacity[string](len(identityPrincipals))
for _, principal := range identityPrincipals {
if allowed.Contains(principal) {
result.Add(principal)
}
}
return result
}
/*
* 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 (
"net/http"
"sort"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
apiclient "github.com/gravitational/teleport/api/client"
"github.com/gravitational/teleport/api/client/proto"
apidefaults "github.com/gravitational/teleport/api/defaults"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
func (h *Handler) getUserGroups(_ http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
// Get a client to the Auth Server with the logged in user's identity. The
// identity of the logged in user is used to fetch the list of nodes.
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
req, err := convertListResourcesRequest(r, types.KindUserGroup)
if err != nil {
return nil, trace.Wrap(err)
}
page, err := apiclient.GetResourcePage[types.UserGroup](r.Context(), clt, req)
if err != nil {
return nil, trace.Wrap(err)
}
appServers, err := apiclient.GetAllResources[types.AppServer](r.Context(), clt, &proto.ListResourcesRequest{
ResourceType: types.KindAppServer,
Namespace: apidefaults.Namespace,
UseSearchAsRoles: true,
})
if err != nil {
h.logger.DebugContext(r.Context(), "Unable to fetch applications while listing user groups, unable to display associated applications", "error", err)
}
appServerLookup := make(map[string]types.AppServer, len(appServers))
for _, appServer := range appServers {
appServerLookup[appServer.GetApp().GetName()] = appServer
}
userGroupsToApps := map[string]types.Apps{}
for _, userGroup := range page.Resources {
apps := make(types.Apps, 0, len(userGroup.GetApplications()))
for _, appName := range userGroup.GetApplications() {
app := appServerLookup[appName]
if app == nil {
h.logger.DebugContext(r.Context(), "Unable to find application when creating user groups, skipping", "app", appName)
continue
}
apps = append(apps, app.GetApp())
}
sort.Sort(apps)
userGroupsToApps[userGroup.GetName()] = apps
}
userGroups, err := ui.MakeUserGroups(page.Resources, userGroupsToApps)
if err != nil {
return nil, trace.Wrap(err)
}
return listResourcesGetResponse{
Items: userGroups,
StartKey: page.NextKey,
TotalCount: page.Total,
}, nil
}
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
"github.com/gravitational/teleport/lib/httplib"
usagereporter "github.com/gravitational/teleport/lib/usagereporter/web"
)
// createPreUserEventHandle sends a user event to the UserEvent service
// this handler is for on-boarding user events pre-session
func (h *Handler) createPreUserEventHandle(w http.ResponseWriter, r *http.Request, p httprouter.Params) (any, error) {
var req usagereporter.CreatePreUserEventRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
client := h.cfg.ProxyClient
typedEvent, err := usagereporter.ConvertPreUserEventRequestToUsageEvent(req)
if err != nil {
return nil, trace.Wrap(err)
}
event := &proto.SubmitUsageEventRequest{
Event: typedEvent,
}
err = client.SubmitUsageEvent(r.Context(), event)
if err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// createUserEventHandle sends a user event to the UserEvent service
// this handler is for user events with a session
func (h *Handler) createUserEventHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, sctx *SessionContext) (any, error) {
var req usagereporter.CreateUserEventRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
client, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
typedEvent, err := usagereporter.ConvertUserEventRequestToUsageEvent(req)
if err != nil {
return nil, trace.Wrap(err)
}
event := &proto.SubmitUsageEventRequest{
Event: typedEvent,
}
err = client.SubmitUsageEvent(r.Context(), event)
if err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
/*
* 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
userpreferencesv1 "github.com/gravitational/teleport/api/gen/proto/go/userpreferences/v1"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
)
// AssistUserPreferencesResponse is the JSON response for the assist user preferences.
type AssistUserPreferencesResponse struct {
PreferredLogins []string `json:"preferredLogins"`
ViewMode userpreferencesv1.AssistViewMode `json:"viewMode"`
}
type preferencesMarketingParams struct {
Campaign string `json:"campaign"`
Source string `json:"source"`
Medium string `json:"medium"`
Intent string `json:"intent"`
}
type OnboardUserPreferencesResponse struct {
PreferredResources []userpreferencesv1.Resource `json:"preferredResources"`
MarketingParams preferencesMarketingParams `json:"marketingParams"`
}
// ClusterUserPreferencesResponse is the JSON response for the user's cluster preferences.
type ClusterUserPreferencesResponse struct {
PinnedResources []string `json:"pinnedResources"`
}
type UnifiedResourcePreferencesResponse struct {
DefaultTab userpreferencesv1.DefaultTab `json:"defaultTab"`
ViewMode userpreferencesv1.ViewMode `json:"viewMode"`
LabelsViewMode userpreferencesv1.LabelsViewMode `json:"labelsViewMode"`
AvailableResourceMode userpreferencesv1.AvailableResourceMode `json:"availableResourceMode"`
}
// AccessGraphPreferencesResponse is the JSON response for Access Graph preferences.
type AccessGraphPreferencesResponse struct {
HasBeenRedirected bool `json:"hasBeenRedirected"`
}
// DiscoverGuidePreferences defines preferences related to discover guides.
type DiscoverGuidePreferences struct {
// PinnedGuides is a list of ids of pinned guides.
Pinned []string `json:"pinned"`
}
// DiscoverResourcePreferencesResponse is the JSON response for discover resource preference
// as part of the user preference request.
type DiscoverResourcePreferencesResponse struct {
DiscoverGuide *DiscoverGuidePreferences `json:"discoverGuide"`
}
// UserPreferencesResponse is the JSON response for the user preferences.
type UserPreferencesResponse struct {
Assist AssistUserPreferencesResponse `json:"assist"`
Theme userpreferencesv1.Theme `json:"theme"`
UnifiedResourcePreferences UnifiedResourcePreferencesResponse `json:"unifiedResourcePreferences"`
Onboard OnboardUserPreferencesResponse `json:"onboard"`
ClusterPreferences ClusterUserPreferencesResponse `json:"clusterPreferences"`
DiscoverResourcePreferences DiscoverResourcePreferencesResponse `json:"discoverResourcePreferences"`
AccessGraph AccessGraphPreferencesResponse `json:"accessGraph"`
SideNavDrawerMode userpreferencesv1.SideNavDrawerMode `json:"sideNavDrawerMode"`
KeyboardLayout uint32 `json:"keyboardLayout"`
}
func (h *Handler) getUserClusterPreferences(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
authClient, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := authClient.GetUserPreferences(r.Context(), &userpreferencesv1.GetUserPreferencesRequest{})
if err != nil {
return nil, trace.Wrap(err)
}
return clusterPreferencesResponse(resp.GetPreferences().GetClusterPreferences()), nil
}
// updateUserClusterPreferences is a handler for PUT /webapi/user/preferences.
func (h *Handler) updateUserClusterPreferences(_ http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
req := UserPreferencesResponse{}
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
authClient, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
preferences := makePreferenceRequest(req)
if err := authClient.UpsertUserPreferences(r.Context(), preferences); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// getUserPreferences is a handler for GET /webapi/user/preferences.
func (h *Handler) getUserPreferences(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext) (any, error) {
authClient, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
resp, err := authClient.GetUserPreferences(r.Context(), &userpreferencesv1.GetUserPreferencesRequest{})
if err != nil {
return nil, trace.Wrap(err)
}
return userPreferencesResponse(resp.GetPreferences()), nil
}
func makePreferenceRequest(req UserPreferencesResponse) *userpreferencesv1.UpsertUserPreferencesRequest {
var discoverGuide *userpreferencesv1.DiscoverGuide
if req.DiscoverResourcePreferences.DiscoverGuide != nil {
discoverGuide = userpreferencesv1.DiscoverGuide_builder{
Pinned: req.DiscoverResourcePreferences.DiscoverGuide.Pinned,
}.Build()
}
return userpreferencesv1.UpsertUserPreferencesRequest_builder{
Preferences: userpreferencesv1.UserPreferences_builder{
KeyboardLayout: req.KeyboardLayout,
Theme: req.Theme,
UnifiedResourcePreferences: userpreferencesv1.UnifiedResourcePreferences_builder{
DefaultTab: req.UnifiedResourcePreferences.DefaultTab,
ViewMode: req.UnifiedResourcePreferences.ViewMode,
LabelsViewMode: req.UnifiedResourcePreferences.LabelsViewMode,
AvailableResourceMode: req.UnifiedResourcePreferences.AvailableResourceMode,
}.Build(),
Onboard: userpreferencesv1.OnboardUserPreferences_builder{
PreferredResources: req.Onboard.PreferredResources,
MarketingParams: userpreferencesv1.MarketingParams_builder{
Campaign: req.Onboard.MarketingParams.Campaign,
Source: req.Onboard.MarketingParams.Source,
Medium: req.Onboard.MarketingParams.Medium,
Intent: req.Onboard.MarketingParams.Intent,
}.Build(),
}.Build(),
ClusterPreferences: userpreferencesv1.ClusterUserPreferences_builder{
PinnedResources: userpreferencesv1.PinnedResourcesUserPreferences_builder{
ResourceIds: req.ClusterPreferences.PinnedResources,
}.Build(),
}.Build(),
AccessGraph: userpreferencesv1.AccessGraphUserPreferences_builder{
HasBeenRedirected: req.AccessGraph.HasBeenRedirected,
}.Build(),
SideNavDrawerMode: req.SideNavDrawerMode,
DiscoverResourcePreferences: userpreferencesv1.DiscoverResourcePreferences_builder{
DiscoverGuide: discoverGuide,
}.Build(),
}.Build(),
}.Build()
}
// updateUserPreferences is a handler for PUT /webapi/user/preferences.
func (h *Handler) updateUserPreferences(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext) (any, error) {
var req UserPreferencesResponse
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
authClient, err := sctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
preferences := makePreferenceRequest(req)
if err := authClient.UpsertUserPreferences(r.Context(), preferences); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
// userPreferencesResponse creates a JSON response for the user preferences.
func userPreferencesResponse(resp *userpreferencesv1.UserPreferences) *UserPreferencesResponse {
jsonResp := &UserPreferencesResponse{
Theme: resp.GetTheme(),
Onboard: onboardUserPreferencesResponse(resp.GetOnboard()),
ClusterPreferences: clusterPreferencesResponse(resp.GetClusterPreferences()),
UnifiedResourcePreferences: unifiedResourcePreferencesResponse(resp.GetUnifiedResourcePreferences()),
AccessGraph: accessGraphPreferencesResponse(resp.GetAccessGraph()),
SideNavDrawerMode: resp.GetSideNavDrawerMode(),
DiscoverResourcePreferences: discoverResourcePreferenceResponse(resp.GetDiscoverResourcePreferences()),
KeyboardLayout: resp.GetKeyboardLayout(),
}
return jsonResp
}
func clusterPreferencesResponse(prefs *userpreferencesv1.ClusterUserPreferences) ClusterUserPreferencesResponse {
resp := ClusterUserPreferencesResponse{}
if prefs == nil {
return resp
}
resp.PinnedResources = append(resp.PinnedResources, prefs.GetPinnedResources().GetResourceIds()...)
return resp
}
// unifiedResourcePreferencesResponse creates a JSON response for the assist user preferences.
func unifiedResourcePreferencesResponse(resp *userpreferencesv1.UnifiedResourcePreferences) UnifiedResourcePreferencesResponse {
return UnifiedResourcePreferencesResponse{
DefaultTab: resp.GetDefaultTab(),
ViewMode: resp.GetViewMode(),
LabelsViewMode: resp.GetLabelsViewMode(),
AvailableResourceMode: resp.GetAvailableResourceMode(),
}
}
// onboardUserPreferencesResponse creates a JSON response for the onboard user preferences.
func onboardUserPreferencesResponse(resp *userpreferencesv1.OnboardUserPreferences) OnboardUserPreferencesResponse {
jsonResp := OnboardUserPreferencesResponse{
PreferredResources: make([]userpreferencesv1.Resource, 0, len(resp.GetPreferredResources())),
MarketingParams: preferencesMarketingParams{
Campaign: resp.GetMarketingParams().GetCampaign(),
Source: resp.GetMarketingParams().GetSource(),
Medium: resp.GetMarketingParams().GetMedium(),
Intent: resp.GetMarketingParams().GetIntent(),
},
}
jsonResp.PreferredResources = append(jsonResp.PreferredResources, resp.GetPreferredResources()...)
return jsonResp
}
// accessGraphPreferencesResponse creates a JSON response for the access graph preferences.
func accessGraphPreferencesResponse(resp *userpreferencesv1.AccessGraphUserPreferences) AccessGraphPreferencesResponse {
if resp == nil {
return AccessGraphPreferencesResponse{
HasBeenRedirected: false,
}
}
return AccessGraphPreferencesResponse{
HasBeenRedirected: resp.GetHasBeenRedirected(),
}
}
// discoverResourcePreferenceResponse creates a JSON response for the discover resource preferences.
func discoverResourcePreferenceResponse(resp *userpreferencesv1.DiscoverResourcePreferences) DiscoverResourcePreferencesResponse {
if resp == nil || resp.GetDiscoverGuide() == nil {
return DiscoverResourcePreferencesResponse{}
}
return DiscoverResourcePreferencesResponse{
DiscoverGuide: &DiscoverGuidePreferences{
Pinned: resp.GetDiscoverGuide().GetPinned(),
},
}
}
/*
* 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"
"net/http"
"time"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
"github.com/gravitational/teleport/api/client/proto"
userspb "github.com/gravitational/teleport/api/gen/proto/go/teleport/users/v1"
"github.com/gravitational/teleport/api/mfa"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/api/utils/clientutils"
"github.com/gravitational/teleport/lib/client"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/web/ui"
)
func (h *Handler) updateUserHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
return updateUser(r, clt)
}
func (h *Handler) createUserHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
return createUser(r, clt, ctx.GetUser())
}
// TODO(rudream): DELETE IN V21.0.0
func (h *Handler) getUsersHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
return getUsers(r.Context(), clt)
}
// listUsersHandle returns a paginated list of users.
func (h *Handler) listUsersHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
users, nextToken, _, err := clientutils.Page(
r.Context(),
int(limit),
values.Get("startKey"),
func(ctx context.Context, pageSize int, pageToken string) ([]*types.UserV2, string, error) {
resp, err := clt.ListUsers(r.Context(), userspb.ListUsersRequest_builder{
PageSize: int32(pageSize),
PageToken: pageToken,
Filter: &types.UserFilter{
SearchKeywords: client.ParseSearchKeywords(values.Get("search"), ' '),
SkipSystemUsers: true,
},
}.Build())
if err != nil {
return nil, "", trace.Wrap(err)
}
return resp.GetUsers(), resp.GetNextPageToken(), nil
})
if err != nil {
return nil, trace.Wrap(err)
}
var uiUsers []ui.UserListEntry
for _, u := range users {
uiuser, err := ui.NewUserListEntry(u)
if err != nil {
return nil, trace.Wrap(err)
}
uiUsers = append(uiUsers, *uiuser)
}
return &listUsersResponse{
Items: uiUsers,
StartKey: nextToken,
}, nil
}
func (h *Handler) getUserHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
username := params.ByName("username")
if username == "" {
return nil, trace.BadParameter("missing username")
}
return getUser(r.Context(), username, clt)
}
func (h *Handler) deleteUserHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
if err := deleteUser(r, params, clt, ctx.GetUser()); err != nil {
return nil, trace.Wrap(err)
}
return OK(), nil
}
func createUser(r *http.Request, m userAPIGetter, createdBy string) (*ui.User, error) {
var req *saveUserRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
user, err := types.NewUser(req.Name)
if err != nil {
return nil, trace.Wrap(err)
}
user.SetRoles(req.Roles)
// checkAndSetDefaults makes sure either TraitsPreset
// or AllTraits field to be populated. Since empty
// AllTraits is also used to delete all user traits,
// we explicitly check if TraitsPreset is empty so
// to prevent traits deletion.
if req.TraitsPreset == nil {
user.SetTraits(req.AllTraits)
} else {
updateUserTraitsPreset(req, user)
}
user.SetCreatedBy(types.CreatedBy{
User: types.UserRef{Name: createdBy},
Time: time.Now().UTC(),
})
created, err := m.CreateUser(r.Context(), user)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.NewUser(created)
}
// updateUserTraitsPreset receives a saveUserRequest and updates the user traits
// accordingly. It only updates the traits that have a non-nil value in
// saveUserRequest. This allows the partial update of the properties
func updateUserTraitsPreset(req *saveUserRequest, user types.User) {
if req.TraitsPreset.Logins != nil {
user.SetLogins(*req.TraitsPreset.Logins)
}
if req.TraitsPreset.DatabaseUsers != nil {
user.SetDatabaseUsers(*req.TraitsPreset.DatabaseUsers)
}
if req.TraitsPreset.DatabaseNames != nil {
user.SetDatabaseNames(*req.TraitsPreset.DatabaseNames)
}
if req.TraitsPreset.KubeUsers != nil {
user.SetKubeUsers(*req.TraitsPreset.KubeUsers)
}
if req.TraitsPreset.KubeGroups != nil {
user.SetKubeGroups(*req.TraitsPreset.KubeGroups)
}
if req.TraitsPreset.WindowsLogins != nil {
user.SetWindowsLogins(*req.TraitsPreset.WindowsLogins)
}
if req.TraitsPreset.AWSRoleARNs != nil {
user.SetAWSRoleARNs(*req.TraitsPreset.AWSRoleARNs)
}
}
func updateUser(r *http.Request, m userAPIGetter) (*ui.User, error) {
var req *saveUserRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.checkAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
// Remove the MFA resp from the context before getting the user.
// Otherwise, it will be consumed before the Update which actually
// requires the MFA.
// TODO(Joerger): Explicitly provide MFA response only where it is
// needed instead of removing it like this.
getUserCtx := mfa.ContextWithMFAResponse(r.Context(), nil)
user, err := m.GetUser(getUserCtx, req.Name, false)
if err != nil {
return nil, trace.Wrap(err)
}
user.SetRoles(req.Roles)
// checkAndSetDefaults makes sure either TraitsPreset
// or AllTraits field to be populated. Since empty
// AllTraits is also used to delete all user traits,
// we explicitly check if TraitsPreset is empty so
// to prevent traits deletion.
if req.TraitsPreset == nil {
user.SetTraits(req.AllTraits)
} else {
updateUserTraitsPreset(req, user)
}
updated, err := m.UpdateUser(r.Context(), user)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.NewUser(updated)
}
func getUsers(ctx context.Context, m userAPIGetter) ([]ui.UserListEntry, error) {
users, err := m.GetUsers(ctx, false)
if err != nil {
return nil, trace.Wrap(err)
}
var uiUsers []ui.UserListEntry
for _, u := range users {
// Do not display system users in the WebUI
if types.IsSystemResource(u) {
continue
}
uiuser, err := ui.NewUserListEntry(u)
if err != nil {
return nil, trace.Wrap(err)
}
uiUsers = append(uiUsers, *uiuser)
}
return uiUsers, nil
}
// listUsersResponse is the response for the list users request.
type listUsersResponse struct {
// Items is the list of users retrieved.
Items []ui.UserListEntry `json:"items"`
// StartKey is the position from which to resume search.
StartKey string `json:"startKey"`
}
func getUser(ctx context.Context, username string, m userAPIGetter) (*ui.User, error) {
user, err := m.GetUser(ctx, username, false)
if err != nil {
return nil, trace.Wrap(err)
}
uiUser, err := ui.NewUser(user)
if err != nil {
return nil, trace.Wrap(err)
}
return uiUser, nil
}
func deleteUser(r *http.Request, params httprouter.Params, m userAPIGetter, user string) error {
username := params.ByName("username")
if username == "" {
return trace.BadParameter("missing user name")
}
if username == user {
return trace.BadParameter("cannot delete own user account")
}
if err := m.DeleteUser(r.Context(), username); err != nil {
return trace.Wrap(err)
}
return nil
}
type privilegeTokenRequest struct {
// ExistingMFAResponse is an MFA challenge response from an existing device.
// Not required if the user has no existing devices.
ExistingMFAResponse *client.MFAChallengeResponse `json:"existingMfaResponse"`
}
// createPrivilegeTokenHandle creates and returns a privilege token.
func (h *Handler) createPrivilegeTokenHandle(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
var req privilegeTokenRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
mfaResp, err := req.ExistingMFAResponse.GetOptionalMFAResponseProtoReq()
if err != nil {
return nil, trace.Wrap(err)
}
clt, err := ctx.GetClient()
if err != nil {
return nil, trace.Wrap(err)
}
token, err := clt.CreatePrivilegeToken(r.Context(), &proto.CreatePrivilegeTokenRequest{ExistingMFAResponse: mfaResp})
if err != nil {
return nil, trace.Wrap(err)
}
return token.GetName(), nil
}
type userAPIGetter interface {
// GetUser returns user by name
GetUser(ctx context.Context, name string, withSecrets bool) (types.User, error)
// CreateUser creates a new user
CreateUser(ctx context.Context, user types.User) (types.User, error)
// UpdateUser updates a user
UpdateUser(ctx context.Context, user types.User) (types.User, error)
// GetUsers returns a list of all users
// TODO(rudream): DELETE IN V21.0.0
GetUsers(ctx context.Context, withSecrets bool) ([]types.User, error)
// ListUsers returns a paginated list of users.
ListUsers(ctx context.Context, req *userspb.ListUsersRequest) (*userspb.ListUsersResponse, error)
// DeleteUser deletes a user by name.
DeleteUser(ctx context.Context, user string) error
}
// traitsPreset are user traits that are pre-defined in Teleport
type traitsPreset struct {
Logins *[]string `json:"logins,omitempty"`
DatabaseUsers *[]string `json:"databaseUsers,omitempty"`
DatabaseNames *[]string `json:"databaseNames,omitempty"`
KubeUsers *[]string `json:"kubeUsers,omitempty"`
KubeGroups *[]string `json:"kubeGroups,omitempty"`
WindowsLogins *[]string `json:"windowsLogins,omitempty"`
AWSRoleARNs *[]string `json:"awsRoleArns,omitempty"`
}
// saveUserRequest represents a create/update request for a user
// Name and Roles are always required
// The remaining fields are part of the Trait map
// They are optional and respect the following logic:
// - if the value is nil, we ignore it
// - if the value is an empty array we remove every element from the trait
// - otherwise, we replace the list for that trait.
// Use TraitsPreset to selectively update traits.
// Use AllTraits to fully replace existing traits.
type saveUserRequest struct {
// Name is username.
Name string `json:"name"`
// Roles is slice of user roles assigned to user.
Roles []string `json:"roles"`
// TraitsPreset holds traits that are pre-defined in Teleport.
// Clients may use TraitsPreset to selectively update user traits.
TraitsPreset *traitsPreset `json:"traits"`
// AllTraits may hold all the user traits, including traits key defined
// in TraitsPreset and/or new trait key values defined by Teleport admin.
// AllTraits should be used to fully replace and update user traits.
AllTraits map[string][]string `json:"allTraits"`
}
func (r *saveUserRequest) checkAndSetDefaults() error {
if r.Name == "" {
return trace.BadParameter("missing user name")
}
if len(r.Roles) == 0 {
return trace.BadParameter("missing roles")
}
if len(r.AllTraits) != 0 && r.TraitsPreset != nil {
return trace.BadParameter("either traits or allTraits must be provided")
}
return nil
}
/*
* Teleport
* Copyright (C) 2024 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 (
"net/http"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
usertasksv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/usertasks/v1"
"github.com/gravitational/teleport/lib/defaults"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/reversetunnelclient"
"github.com/gravitational/teleport/lib/web/ui"
)
// userTaskStateUpdate updates the state of a User Task.
func (h *Handler) userTaskStateUpdate(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
userTaskName := p.ByName("name")
if userTaskName == "" {
return nil, trace.BadParameter("a user task name is required")
}
var req *ui.UpdateUserTaskStateRequest
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.CheckAndSetDefaults(); err != nil {
return nil, trace.Wrap(err)
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
userTask, err := clt.UserTasksServiceClient().GetUserTask(r.Context(), userTaskName)
if err != nil {
return nil, trace.Wrap(err)
}
userTask.GetSpec().SetState(req.State)
newUserTask, err := clt.UserTasksServiceClient().UpsertUserTask(r.Context(), userTask)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeUserTask(newUserTask), nil
}
// userTaskGet returns a User Task based on its name
func (h *Handler) userTaskGet(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
userTaskName := p.ByName("name")
if userTaskName == "" {
return nil, trace.BadParameter("a user task name is required")
}
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
ut, err := clt.UserTasksServiceClient().GetUserTask(r.Context(), userTaskName)
if err != nil {
return nil, trace.Wrap(err)
}
return ui.MakeDetailedUserTask(ut), nil
}
// userTaskListByIntegration returns a page of User Tasks.
// It requires a query param to filter by integration
//
// The following query params are optional:
// - limit: max number of items
// - startKey: used to iterate over pages
//
// It returns a list of user tasks with the base attributes (common among all user tasks).
// To get a detailed UserTask use the single resource endpoint, ie, usertask/<resource's name>.
func (h *Handler) userTaskListByIntegration(w http.ResponseWriter, r *http.Request, p httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
values := r.URL.Query()
limit, err := QueryLimitAsInt32(values, "limit", defaults.MaxIterationLimit)
if err != nil {
return nil, trace.Wrap(err)
}
startKey := values.Get("startKey")
integrationName := values.Get("integration")
if integrationName == "" {
return nil, trace.BadParameter("integration query param is required")
}
taskStateFilter := values.Get("state")
filters := usertasksv1.ListUserTasksFilters_builder{
Integration: integrationName,
TaskState: taskStateFilter,
}.Build()
userTasks, nextKey, err := clt.UserTasksServiceClient().ListUserTasks(r.Context(), int64(limit), startKey, filters)
if err != nil {
return nil, trace.Wrap(err)
}
items := make([]ui.UserTask, 0, len(userTasks))
for _, userTask := range userTasks {
items = append(items, ui.MakeUserTask(userTask))
}
return ui.UserTasksListResponse{
Items: items,
NextKey: nextKey,
}, nil
}
/*
* Teleport
* Copyright (C) 2025 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 (
"net/http"
"strconv"
"strings"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
scopesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/scopes/v1"
workloadidentityv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1"
"github.com/gravitational/teleport/lib/reversetunnelclient"
tslices "github.com/gravitational/teleport/lib/utils/slices"
)
// listWorkloadIdentities returns a list of workload identities for a given
// cluster site.
func (h *Handler) listWorkloadIdentities(_ http.ResponseWriter, r *http.Request, _ httprouter.Params, sctx *SessionContext, cluster reversetunnelclient.Cluster) (any, error) {
clt, err := sctx.GetUserClient(r.Context(), cluster)
if err != nil {
return nil, trace.Wrap(err)
}
request := workloadidentityv1.ListWorkloadIdentitiesV2Request_builder{
PageSize: 20,
PageToken: r.URL.Query().Get("page_token"),
SortField: r.URL.Query().Get("sort_field"),
FilterSearchTerm: r.URL.Query().Get("search"),
// Exhaustive view, so ask for every scope rather than inheriting the
// identity-based default.
ScopeFilter: scopesv1.Filter_builder{Mode: scopesv1.Mode_MODE_ALL}.Build(),
}.Build()
if r.URL.Query().Has("page_size") {
pageSize, err := strconv.ParseInt(r.URL.Query().Get("page_size"), 10, 32)
if err != nil {
return nil, trace.BadParameter("invalid page size")
}
request.SetPageSize(int32(pageSize))
}
if r.URL.Query().Has("sort_dir") {
sortDir := r.URL.Query().Get("sort_dir")
request.SetSortDesc(strings.ToLower(sortDir) == "desc")
}
result, err := clt.WorkloadIdentityResourceServiceClient().ListWorkloadIdentitiesV2(r.Context(), request)
if err != nil {
return nil, trace.Wrap(err)
}
uiItems := tslices.Map(result.GetWorkloadIdentities(), func(item *workloadidentityv1.WorkloadIdentity) WorkloadIdentity {
uiItem := WorkloadIdentity{
Name: item.GetMetadata().GetName(),
Scope: item.GetScope(),
SpiffeID: item.GetSpec().GetSpiffe().GetId(),
SpiffeHint: item.GetSpec().GetSpiffe().GetHint(),
Labels: item.GetMetadata().GetLabels(),
}
return uiItem
})
return ListWorkloadIdentitiesResponse{
Items: uiItems,
NextPageToken: result.GetNextPageToken(),
}, nil
}
type ListWorkloadIdentitiesResponse struct {
Items []WorkloadIdentity `json:"items"`
NextPageToken string `json:"next_page_token,omitempty"`
}
type WorkloadIdentity struct {
Name string `json:"name"`
Scope string `json:"scope,omitempty"`
SpiffeID string `json:"spiffe_id"`
SpiffeHint string `json:"spiffe_hint"`
Labels map[string]string `json:"labels"`
}
/*
* 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"
"log/slog"
"time"
"github.com/gorilla/websocket"
"github.com/gravitational/trace"
)
type WebsocketIO struct {
Conn *websocket.Conn
remaining []byte
}
func (ws *WebsocketIO) Write(p []byte) (int, error) {
err := ws.Conn.WriteMessage(websocket.BinaryMessage, p)
if err != nil {
return 0, trace.Wrap(err)
}
return len(p), nil
}
func (ws *WebsocketIO) Read(p []byte) (int, error) {
if len(ws.remaining) == 0 {
ty, data, err := ws.Conn.ReadMessage()
if err != nil {
return 0, trace.Wrap(err)
}
if ty != websocket.BinaryMessage {
return 0, trace.BadParameter("expected websocket message of type BinaryMessage, got %T", ty)
}
ws.remaining = data
}
copied := copy(p, ws.remaining)
ws.remaining = ws.remaining[copied:]
return copied, nil
}
func (ws *WebsocketIO) Close() error {
return trace.Wrap(ws.Conn.Close())
}
type wsPinger interface {
WriteControl(messageType int, data []byte, deadline time.Time) error
}
// startWSPingLoop starts a loop that will continuously send a ping frame through the websocket
// to prevent the connection between web client and teleport proxy from becoming idle.
// Interval is determined by the keep_alive_interval config set by user (or default).
// Loop will terminate when there is an error sending ping frame or when the context is canceled.
func startWSPingLoop(ctx context.Context, pinger wsPinger, keepAliveInterval time.Duration, log *slog.Logger, onClose func() error) {
log.DebugContext(ctx, "Starting websocket ping loop with interval", "interval", keepAliveInterval)
tickerCh := time.NewTicker(keepAliveInterval)
defer tickerCh.Stop()
for {
select {
case <-tickerCh.C:
// A short deadline is used here to detect a broken connection quickly.
// If this is just a temporary issue, we will retry shortly anyway.
deadline := time.Now().Add(time.Second)
if err := pinger.WriteControl(websocket.PingMessage, nil, deadline); err != nil {
log.ErrorContext(ctx, "Unable to send ping frame to web client", "error", err)
if onClose != nil {
if err := onClose(); err != nil {
log.ErrorContext(ctx, "OnClose handler failed", "error", err)
}
}
return
}
case <-ctx.Done():
log.DebugContext(ctx, "Terminating websocket ping loop.")
return
}
}
}
/*
* 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 (
"net/http"
yaml "github.com/ghodss/yaml"
"github.com/gravitational/trace"
"github.com/julienschmidt/httprouter"
accessmonitoringrulesv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/accessmonitoringrules/v1"
"github.com/gravitational/teleport/api/types"
"github.com/gravitational/teleport/lib/httplib"
"github.com/gravitational/teleport/lib/services"
)
type yamlParseRequest struct {
YAML string `json:"yaml"`
}
type yamlParseResponse struct {
Resource any `json:"resource"`
}
type yamlStringifyResponse struct {
YAML string `json:"yaml"`
}
func (h *Handler) yamlParse(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
kind := params.ByName("kind")
if len(kind) == 0 {
return nil, trace.BadParameter("query param %q is required", "kind")
}
var req yamlParseRequest
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
switch kind {
case types.KindAccessMonitoringRule:
resource, err := yamlToAccessMonitoringRuleResource(req.YAML)
if err != nil {
return nil, trace.Wrap(err)
}
return yamlParseResponse{Resource: resource}, nil
case types.KindRole:
resource, err := yamlToRole(req.YAML)
if err != nil {
return nil, trace.Wrap(err)
}
return yamlParseResponse{Resource: resource}, nil
case types.KindToken:
resource, err := yamlToProvisionToken(req.YAML)
if err != nil {
return nil, trace.Wrap(err)
}
return yamlParseResponse{Resource: resource}, nil
default:
return nil, trace.NotImplemented("parsing YAML for kind %q is not supported", kind)
}
}
func (h *Handler) yamlStringify(w http.ResponseWriter, r *http.Request, params httprouter.Params, ctx *SessionContext) (any, error) {
kind := params.ByName("kind")
if len(kind) == 0 {
return nil, trace.BadParameter("query param %q is required", "kind")
}
var resource any
switch kind {
case types.KindAccessMonitoringRule:
var req struct {
Resource *accessmonitoringrulesv1.AccessMonitoringRule `json:"resource"`
}
if err := httplib.ReadResourceJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
resource = req.Resource
case types.KindRole:
var req struct {
Resource types.RoleV6 `json:"resource"`
}
if err := httplib.ReadJSON(r, &req); err != nil {
return nil, trace.Wrap(err)
}
if err := req.Resource.CheckAndSetDefaults(); err != nil {
return nil, err
}
resource = req.Resource
default:
return nil, trace.NotImplemented("YAML stringifying for kind %q is not supported", kind)
}
data, err := yaml.Marshal(resource)
if err != nil {
return nil, trace.Wrap(err)
}
return yamlStringifyResponse{YAML: string(data)}, nil
}
func yamlToAccessMonitoringRuleResource(yaml string) (*accessmonitoringrulesv1.AccessMonitoringRule, error) {
extractedRes, err := extractResource(yaml)
if err != nil {
return nil, trace.Wrap(err)
}
if extractedRes.Kind != types.KindAccessMonitoringRule {
return nil, trace.BadParameter("resource kind %q is invalid, only access_monitoring_rule is allowed", extractedRes.Kind)
}
resource, err := services.UnmarshalAccessMonitoringRule(extractedRes.Raw)
if err != nil {
return nil, trace.Wrap(err)
}
return resource, nil
}
func yamlToRole(yaml string) (types.Role, error) {
extractedRes, err := extractResource(yaml)
if err != nil {
return nil, trace.Wrap(err)
}
if extractedRes.Kind != types.KindRole {
return nil, trace.BadParameter("resource kind %q is invalid, only role is allowed", extractedRes.Kind)
}
resource, err := services.UnmarshalRole(extractedRes.Raw, services.DisallowUnknown())
if err != nil {
return nil, trace.Wrap(err)
}
return resource, nil
}
func yamlToProvisionToken(yaml string) (types.ProvisionToken, error) {
extractedRes, err := extractResource(yaml)
if err != nil {
return nil, trace.Wrap(err)
}
if extractedRes.Kind != types.KindToken {
return nil, trace.BadParameter("resource kind %q is invalid, only token is allowed", extractedRes.Kind)
}
resource, err := services.UnmarshalProvisionToken(extractedRes.Raw)
if err != nil {
return nil, trace.Wrap(err)
}
return resource, nil
}
/*
* 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 x11
import (
"context"
"crypto/rand"
"encoding/binary"
"encoding/hex"
"fmt"
"io"
"os/exec"
"strings"
"time"
"github.com/gravitational/trace"
)
const (
// XAuthFileEnvVar is the environment variable used to specify the path to the
// X11 authorization file.
XAuthFileEnvVar = "XAUTHORITY"
// mitMagicCookieProto is the default xauth protocol used for X11 forwarding.
mitMagicCookieProto = "MIT-MAGIC-COOKIE-1"
// mitMagicCookieSize is the number of bytes in an mit magic cookie.
mitMagicCookieSize = 16
)
// XAuthEntry is an entry in an XAuthority database which can be used to authenticate
// and authorize requests from an XServer to the associated X display.
type XAuthEntry struct {
// Display is an X display in the format - [hostname]:[display_number].[screen_number]
Display Display `json:"display"`
// Proto is an XAuthority protocol, generally "MIT-MAGIC-COOKIE-1"
Proto string `json:"proto"`
// Cookie is a hex encoded XAuthority cookie
Cookie string `json:"cookie"`
}
// NewFakeXAuthEntry creates a fake xauth entry with a randomly generated MIT-MAGIC-COOKIE-1.
func NewFakeXAuthEntry(display Display) (*XAuthEntry, error) {
cookie, err := newCookie(mitMagicCookieSize)
if err != nil {
return nil, trace.Wrap(err)
}
return &XAuthEntry{
Display: display,
Proto: mitMagicCookieProto,
Cookie: cookie,
}, nil
}
// SpoofXAuthEntry creates a new xauth entry with a random cookie with the
// same length as the original entry's cookie. This is used to create a
// believable spoof of the client's xauth data to send to the server.
func (e *XAuthEntry) SpoofXAuthEntry() (*XAuthEntry, error) {
spoofedCookie, err := newCookie(hex.DecodedLen(len(e.Cookie)))
if err != nil {
return nil, trace.Wrap(err)
}
return &XAuthEntry{
Display: e.Display,
Proto: e.Proto,
Cookie: spoofedCookie,
}, nil
}
// newCookie makes a random hex-encoded cookie with the given byte length.
func newCookie(byteLength int) (string, error) {
cookieBytes := make([]byte, byteLength)
if _, err := rand.Read(cookieBytes); err != nil {
return "", trace.Wrap(err)
}
return hex.EncodeToString(cookieBytes), nil
}
// XAuthCommand is a os/exec.Cmd wrapper for running xauth commands.
type XAuthCommand struct {
*exec.Cmd
}
// NewXAuthCommand reate a new "xauth" command. xauthFile can be
// optionally provided to run the xauth command against a specific xauth file.
func NewXAuthCommand(ctx context.Context, xauthFile string) *XAuthCommand {
var args []string
if xauthFile != "" {
args = []string{"-f", xauthFile}
}
return &XAuthCommand{exec.CommandContext(ctx, "xauth", args...)}
}
// ReadEntry runs "xauth list" to read the first xauth entry for the given display.
func (x *XAuthCommand) ReadEntry(display Display) (*XAuthEntry, error) {
x.Cmd.Args = append(x.Cmd.Args, "list", display.String())
out, err := x.output()
if err != nil {
return nil, trace.Wrap(err)
}
if len(out) == 0 {
return nil, trace.NotFound("no xauth entry found")
}
// Ignore entries beyond the first listed.
entry := strings.Split(string(out), "\n")[0]
splitEntry := strings.Split(entry, " ")
if len(splitEntry) != 3 {
return nil, trace.Errorf("invalid xAuthEntry, expected entry to have three parts")
}
proto, cookie := splitEntry[1], splitEntry[2]
return &XAuthEntry{
Display: display,
Proto: proto,
Cookie: cookie,
}, nil
}
// RemoveEntries runs "xauth remove" to remove any xauth entries for the given display.
func (x *XAuthCommand) RemoveEntries(display Display) error {
x.Cmd.Args = append(x.Cmd.Args, "remove", display.String())
return trace.Wrap(x.run())
}
// AddEntry runs "xauth add" to add the given xauth entry.
func (x *XAuthCommand) AddEntry(entry XAuthEntry) error {
x.Cmd.Args = append(x.Cmd.Args, "add", entry.Display.String(), entry.Proto, entry.Cookie)
return trace.Wrap(x.run())
}
// GenerateUntrustedCookie runs "xauth generate untrusted" to create a new xauth entry with
// an untrusted MIT-MAGIC-COOKIE-1. A timeout can optionally be set for the xauth entry, after
// which the XServer will ignore this cookie.
func (x *XAuthCommand) GenerateUntrustedCookie(display Display, timeout time.Duration) error {
x.Cmd.Args = append(x.Cmd.Args, "generate", display.String(), mitMagicCookieProto, "untrusted")
x.Cmd.Args = append(x.Cmd.Args, "timeout", fmt.Sprint(timeout/time.Second))
return trace.Wrap(x.run())
}
// run the command and return stderr if there is an error.
func (x *XAuthCommand) run() error {
_, err := x.output()
return trace.Wrap(err)
}
// run the command and return stdout or stderr if there is an error.
func (x *XAuthCommand) output() ([]byte, error) {
stdout, err := x.Cmd.StdoutPipe()
if err != nil {
return nil, trace.Wrap(err)
}
stderr, err := x.Cmd.StderrPipe()
if err != nil {
return nil, trace.Wrap(err)
}
if err := x.Cmd.Start(); err != nil {
return nil, trace.Wrap(err)
}
// We add a conservative peak length of 10 KB to prevent potential
// output spam from the client provided `xauth` binary
var peakLength int64 = 10000
out, err := io.ReadAll(io.LimitReader(stdout, peakLength))
if err != nil {
return nil, trace.Wrap(err)
}
errOut, err := io.ReadAll(io.LimitReader(stderr, peakLength))
if err != nil {
return nil, trace.Wrap(err)
}
if err := x.Wait(); err != nil {
return nil, trace.Wrap(err, "command \"%s\" failed with stderr: \"%s\"", strings.Join(x.Cmd.Args, " "), errOut)
}
return out, nil
}
// CheckXAuthPath checks if xauth is runnable in the current environment.
func CheckXAuthPath() error {
_, err := exec.LookPath("xauth")
return trace.Wrap(err)
}
// ReadAndRewriteXAuthPacket reads the initial xauth packet from an XServer request. The xauth packet has 2 parts:
// 1. fixed size buffer (12 bytes) - holds byteOrder bit, and the sizes of the protocol string and auth data
// 2. variable size xauth packet - holds xauth protocol and data used to connect to the remote XServer.
//
// Then it compares the received auth packet with the auth proto and fake cookie
// sent to the server with the original "x11-req". If the data matches, the auth
// packet is returned with the fake cookie replaced by the real cookie to provide
// access to the client's X display.
func ReadAndRewriteXAuthPacket(xreq io.Reader, spoofedXAuthEntry, realXAuthEntry *XAuthEntry) ([]byte, error) {
if spoofedXAuthEntry.Proto != realXAuthEntry.Proto || len(spoofedXAuthEntry.Cookie) != len(realXAuthEntry.Cookie) {
return nil, trace.BadParameter("spoofed and real xauth entries must use the same xauth protocol")
}
// xauth packet starts with a fixed sized buffer of 12 bytes
// which is used to size and decode the remaining bytes
initBuf := make([]byte, xauthPacketInitBufSize)
if _, err := io.ReadFull(xreq, initBuf); err != nil {
return nil, trace.Wrap(err, "X11 channel initial packet buffer missing or too short")
}
protoLen, dataLen, err := readXauthPacketInitBuf(initBuf)
if err != nil {
return nil, trace.Wrap(err)
}
// authPacket size is equal to protoLen (rounded up by 4) + dataLen.
// In openssh, the rounding is performed with: (protoLen + 3) & ~3
authPacketSize := protoLen + (4-protoLen%4)%4 + dataLen
authPacket := make([]byte, authPacketSize)
if _, err := io.ReadFull(xreq, authPacket); err != nil {
return nil, trace.Wrap(err, "X11 channel auth packet missing or too short")
}
proto := authPacket[:protoLen]
authData := authPacket[len(authPacket)-dataLen:]
if string(proto) != spoofedXAuthEntry.Proto || hex.EncodeToString(authData) != spoofedXAuthEntry.Cookie {
return nil, trace.AccessDenied("X11 channel has the wrong authentication data")
}
// Replace auth data with the real auth data
realAuthData, err := hex.DecodeString(realXAuthEntry.Cookie)
if err != nil {
return nil, trace.Wrap(err)
}
copy(authData, realAuthData)
return append(initBuf, authPacket...), trace.Wrap(err)
}
const (
// xauthPacketInitBufSize is the size of the initial
// fixed portion of an xauth packet
xauthPacketInitBufSize = 12
// little endian byte order
littleEndian = 'l'
// big endian byte order
bigEndian = 'B'
)
// readXauthPacketInitBuf reads the initial fixed size portion of
// an xauth packet to get the length of the auth proto and auth data
// portions of the xauth packet.
func readXauthPacketInitBuf(initBuf []byte) (protoLen int, dataLen int, err error) {
// The first byte in the packet determines the
// byte order of the initial buffer's bytes.
var e binary.ByteOrder
switch initBuf[0] {
case bigEndian:
e = binary.BigEndian
case littleEndian:
e = binary.LittleEndian
default:
return 0, 0, trace.BadParameter("X11 channel auth packet has invalid byte order: %v", initBuf[0])
}
// bytes 6-7 and 8-9 are used to determine the length of
// the auth proto and auth data fields respectively.
protoLen = int(e.Uint16(initBuf[6:8]))
dataLen = int(e.Uint16(initBuf[8:10]))
return protoLen, dataLen, nil
}
/*
* 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 x11
import (
"fmt"
"net"
"os"
"path/filepath"
"strconv"
"strings"
"unicode"
"github.com/gravitational/trace"
)
const (
// DisplayEnv is an environment variable used to determine what
// local display should be connected to during X11 forwarding.
DisplayEnv = "DISPLAY"
// x11SocketDirName is the name of the directory where X11 unix sockets are kept.
x11SocketDirName = ".X11-unix"
// x11BasePort is the base port used for XServer tcp addresses.
// e.g. DISPLAY=localhost:10 -> net.Dial("tcp", "localhost:6010")
// Used by some XServer clients, such as openSSH and MobaXTerm.
x11BasePort = 6000
)
// Display is an XServer display.
type Display struct {
// HostName is the display's hostname. For tcp display sockets, this will be
// an ip address. For unix display sockets, this will be empty or "unix".
HostName string `json:"hostname"`
// DisplayNumber is a number representing an x display.
DisplayNumber int `json:"display_number"`
// ScreenNumber is a specific screen number of an x display.
ScreenNumber int `json:"screen_number"`
}
// String returns the string representation of the display.
func (d *Display) String() string {
return fmt.Sprintf("%s:%d.%d", d.HostName, d.DisplayNumber, d.ScreenNumber)
}
func (d *Display) getNetAddr() (net.Addr, error) {
unixSock, unixErr := d.unixSocket()
if unixErr == nil {
return unixSock, nil
}
tcpSock, tcpErr := d.tcpSocket()
if tcpErr == nil {
return tcpSock, nil
}
return nil, trace.NewAggregate(unixErr, tcpErr)
}
// Dial opens an XServer connection to the display
func (d *Display) Dial() (net.Conn, error) {
netAddr, err := d.getNetAddr()
if err != nil {
return nil, trace.Wrap(err)
}
conn, err := net.Dial(netAddr.Network(), netAddr.String())
return conn, trace.Wrap(err)
}
// Listen opens an XServer listener. It will attempt to listen on the display
// address for both tcp and unix and return an aggregate error, unless one
// results in an addr in use error.
func (d *Display) Listen() (net.Listener, error) {
netAddr, err := d.getNetAddr()
if err != nil {
return nil, trace.Wrap(err)
}
listener, err := net.Listen(netAddr.Network(), netAddr.String())
return listener, trace.Wrap(err)
}
// xserverUnixSocket returns the display's associated unix socket.
func (d *Display) unixSocket() (*net.UnixAddr, error) {
// If hostname is "unix" or empty, then the actual unix socket
// for the display is "/tmp/.X11-unix/X<display_number>"
if d.HostName == "unix" || d.HostName == "" {
sockName := filepath.Join(x11SockDir(), fmt.Sprintf("X%d", d.DisplayNumber))
return net.ResolveUnixAddr("unix", sockName)
}
// It's possible that the display is actually the full path
// to an open XServer socket, such as with xquartz on OSX:
// "/private/tmp/com.apple.com/launchd.xxx/org.xquartz.com:0"
if d.HostName[0] == '/' {
sockName := d.String()
if _, err := os.Stat(sockName); err == nil {
return net.ResolveUnixAddr("unix", sockName)
}
// The socket might not include the screen number.
sockName = fmt.Sprintf("%s:%d", d.HostName, d.DisplayNumber)
if _, err := os.Stat(sockName); err == nil {
return net.ResolveUnixAddr("unix", sockName)
}
}
return nil, trace.BadParameter("display is not a unix socket")
}
// ParseDisplay parses the given display value and returns the host,
// display number, and screen number, or a parsing error. display must be
// in one of the following formats - hostname:d[.s], unix:d[.s], :d[.s], ::d[.s].
func ParseDisplayFromUnixSocket(socket string) (Display, error) {
if filepath.Dir(socket) != x11SockDir() {
return Display{}, trace.BadParameter("parsing x11 sockets outside of the standard /tmp/.X11-unix path is not supported")
}
// The file name should look like X[d].[s] and we want :[d].[s]
fileName := filepath.Base(socket)
displayName := strings.Replace(fileName, "X", ":", 1)
return ParseDisplay(displayName)
}
// xserverTCPSocket returns the display's associated tcp socket.
// e.g. "hostname:<6000+display_number>"
func (d *Display) tcpSocket() (*net.TCPAddr, error) {
if d.HostName == "" {
return nil, trace.BadParameter("display is not a tcp socket, hostname can't be empty")
}
port := fmt.Sprint(d.DisplayNumber + x11BasePort)
rawAddr := net.JoinHostPort(d.HostName, port)
addr, err := net.ResolveTCPAddr("tcp", rawAddr)
if err != nil {
return nil, trace.Wrap(err)
}
return addr, nil
}
// GetXDisplay retrieves and validates the local XServer display
// set in the environment variable $DISPLAY.
func GetXDisplay() (Display, error) {
displayString := os.Getenv(DisplayEnv)
if displayString == "" {
return Display{}, trace.BadParameter("$DISPLAY not set")
}
display, err := ParseDisplay(displayString)
if err != nil {
return Display{}, trace.Wrap(err)
}
return display, nil
}
// ParseDisplay parses the given display value and returns the host,
// display number, and screen number, or a parsing error. display must be
// in one of the following formats - hostname:d[.s], unix:d[.s], :d[.s], ::d[.s].
func ParseDisplay(displayString string) (Display, error) {
if displayString == "" {
return Display{}, trace.BadParameter("display cannot be an empty string")
}
// check the display for illegal characters in case of code injection attempt
allowedSpecialChars := ":/.-_" // chars used for hostname or display delimiters.
for _, c := range displayString {
if !unicode.IsLetter(c) && !unicode.IsNumber(c) && !strings.ContainsRune(allowedSpecialChars, c) {
return Display{}, trace.BadParameter("display contains invalid character %q", c)
}
}
// Parse hostname.
colonIdx := strings.LastIndex(displayString, ":")
if colonIdx == -1 || len(displayString) == colonIdx+1 {
return Display{}, trace.BadParameter("display value is missing display number")
}
var display Display
if displayString[0] == ':' {
display.HostName = ""
} else {
display.HostName = displayString[:colonIdx]
}
// Parse display number and screen number
splitDot := strings.Split(displayString[colonIdx+1:], ".")
displayNumber, err := strconv.ParseUint(splitDot[0], 10, 0)
if err != nil {
return Display{}, trace.Wrap(err)
}
display.DisplayNumber = int(displayNumber)
if len(splitDot) < 2 {
return display, nil
}
screenNumber, err := strconv.ParseUint(splitDot[1], 10, 0)
if err != nil {
return Display{}, trace.Wrap(err)
}
display.ScreenNumber = int(screenNumber)
return display, nil
}
func x11SockDir() string {
// We use "/tmp" instead of "os.TempDir" because X11
// is not OS aware and uses "/tmp" regardless of OS.
return filepath.Join("/tmp", x11SocketDirName)
}
/*
* 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 x11
import (
"errors"
"math"
"net"
"os"
"syscall"
"github.com/gravitational/trace"
)
const (
// DefaultDisplayOffset is the default display offset when
// searching for an open XServer unix socket.
DefaultDisplayOffset = 10
// DefaultMaxDisplays is the default maximum number of displays
// supported when searching for an open XServer unix socket.
DefaultMaxDisplays = 1000
// MaxDisplay is the theoretical max display value which
// X Clients and serverwill be able to parse into a unix socket.
MaxDisplayNumber = math.MaxInt32
)
// OpenNewXServerListener opens an XServerListener for the first available Display.
// displayOffset will determine what display number to start from when searching for
// an open display unix socket, and maxDisplays in optional limit for the number of
// display sockets which can be opened at once.
func OpenNewXServerListener(displayOffset int, maxDisplay int, screen uint32) (net.Listener, Display, error) {
if displayOffset > maxDisplay {
return nil, Display{}, trace.BadParameter("displayOffset (%d) cannot be larger than maxDisplay (%d)", displayOffset, maxDisplay)
} else if maxDisplay > MaxDisplayNumber {
return nil, Display{}, trace.BadParameter("maxDisplay (%d) cannot be larger than the max int32 (%d)", maxDisplay, math.MaxInt32)
}
// Create /tmp/.X11-unix if it doesn't exist (such as in CI)
if err := os.Mkdir(x11SockDir(), 0o777|os.ModeSticky); err != nil && !errors.Is(err, os.ErrExist) {
return nil, Display{}, trace.Wrap(err)
}
for displayNumber := displayOffset; displayNumber <= maxDisplay; displayNumber++ {
display := Display{DisplayNumber: displayNumber, ScreenNumber: int(screen)}
if l, err := display.Listen(); err == nil {
return l, display, nil
} else if tryNextDisplayError(err) {
// Skip to next display if the error is non-fatal.
continue
} else {
return nil, Display{}, trace.Wrap(err)
}
}
return nil, Display{}, trace.LimitExceeded("No more X11 sockets are available")
}
// tryNextDisplayError checks if the error is non-fatal and connecting to the
// next display should be attempted. The following cases are supported.
//
// * syscall.EADDRINUSE: Socket exists and something is already bound to it.
// This happens when a display is already being used by another X server.
//
// * syscall.EACCES: Permissions error. For example, socket may be owned by
// different user.
//
// * syscall.EEXIST: Socket already exists. Happens on macOS most often and in
// particular during heavy load on CI.
func tryNextDisplayError(err error) bool {
return errors.Is(err, syscall.EADDRINUSE) ||
errors.Is(err, syscall.EACCES) ||
errors.Is(err, syscall.EEXIST)
}