diff --git a/changes/issue-12613-azure-jwt-support b/changes/issue-12613-azure-jwt-support new file mode 100644 index 0000000000..75c117501c --- /dev/null +++ b/changes/issue-12613-azure-jwt-support @@ -0,0 +1 @@ +* Adding support for Azure JWT tokens diff --git a/changes/issue-12614-adding-support-for-sts-auth-endpoint b/changes/issue-12614-adding-support-for-sts-auth-endpoint new file mode 100644 index 0000000000..fc8b751257 --- /dev/null +++ b/changes/issue-12614-adding-support-for-sts-auth-endpoint @@ -0,0 +1 @@ +* Adding support for Windows MDM STS Auth Endpoint diff --git a/server/datastore/mysql/microsoft_mdm.go b/server/datastore/mysql/microsoft_mdm.go index aea7d99f56..863f9c290f 100644 --- a/server/datastore/mysql/microsoft_mdm.go +++ b/server/datastore/mysql/microsoft_mdm.go @@ -10,7 +10,7 @@ import ( ) // MDMWindowsGetEnrolledDevice receives a Windows MDM device id and returns the device information. -func (ds *Datastore) MDMWindowsGetEnrolledDevice(ctx context.Context, mdmDeviceID string) (*fleet.MDMWindowsEnrolledDevice, error) { +func (ds *Datastore) MDMWindowsGetEnrolledDevice(ctx context.Context, mdmDeviceHWID string) (*fleet.MDMWindowsEnrolledDevice, error) { stmt := `SELECT mdm_device_id, mdm_hardware_id, @@ -24,12 +24,12 @@ func (ds *Datastore) MDMWindowsGetEnrolledDevice(ctx context.Context, mdmDeviceI not_in_oobe, created_at, updated_at - FROM mdm_windows_enrollments WHERE mdm_device_id = ?` + FROM mdm_windows_enrollments WHERE mdm_hardware_id = ?` var winMDMDevice fleet.MDMWindowsEnrolledDevice - if err := sqlx.GetContext(ctx, ds.reader(ctx), &winMDMDevice, stmt, mdmDeviceID); err != nil { + if err := sqlx.GetContext(ctx, ds.reader(ctx), &winMDMDevice, stmt, mdmDeviceHWID); err != nil { if err == sql.ErrNoRows { - return nil, ctxerr.Wrap(ctx, notFound("MDMWindowsEnrolledDevice").WithMessage(mdmDeviceID)) + return nil, ctxerr.Wrap(ctx, notFound("MDMWindowsEnrolledDevice").WithMessage(mdmDeviceHWID)) } return nil, ctxerr.Wrap(ctx, err, "get MDMWindowsEnrolledDevice") } @@ -66,7 +66,7 @@ func (ds *Datastore) MDMWindowsInsertEnrolledDevice(ctx context.Context, device device.MDMNotInOOBE) if err != nil { if isDuplicate(err) { - return ctxerr.Wrap(ctx, alreadyExists("MDMWindowsEnrolledDevice", device.MDMDeviceID)) + return ctxerr.Wrap(ctx, alreadyExists("MDMWindowsEnrolledDevice", device.MDMHardwareID)) } return ctxerr.Wrap(ctx, err, "inserting MDMWindowsEnrolledDevice") } @@ -75,10 +75,10 @@ func (ds *Datastore) MDMWindowsInsertEnrolledDevice(ctx context.Context, device } // MDMWindowsDeleteEnrolledDevice deletes a give MDMWindowsEnrolledDevice entry from the database using the device id. -func (ds *Datastore) MDMWindowsDeleteEnrolledDevice(ctx context.Context, mdmDeviceID string) error { - stmt := "DELETE FROM mdm_windows_enrollments WHERE mdm_device_id = ?" +func (ds *Datastore) MDMWindowsDeleteEnrolledDevice(ctx context.Context, mdmDeviceHWID string) error { + stmt := "DELETE FROM mdm_windows_enrollments WHERE mdm_hardware_id = ?" - res, err := ds.writer(ctx).ExecContext(ctx, stmt, mdmDeviceID) + res, err := ds.writer(ctx).ExecContext(ctx, stmt, mdmDeviceHWID) if err != nil { return ctxerr.Wrap(ctx, err, "delete MDMWindowsEnrolledDevice") } diff --git a/server/datastore/mysql/microsoft_mdm_test.go b/server/datastore/mysql/microsoft_mdm_test.go index d0020e0a86..00e2a15265 100644 --- a/server/datastore/mysql/microsoft_mdm_test.go +++ b/server/datastore/mysql/microsoft_mdm_test.go @@ -33,7 +33,7 @@ func testMDMWindowsEnrolledDevice(t *testing.T, ds *Datastore) { enrolledDevice := &fleet.MDMWindowsEnrolledDevice{ MDMDeviceID: uuid.New().String(), - MDMHardwareID: uuid.New().String(), + MDMHardwareID: uuid.New().String() + uuid.New().String(), MDMDeviceState: uuid.New().String(), MDMDeviceType: "CIMClient_Windows", MDMDeviceName: "DESKTOP-1C3ARC1", @@ -51,19 +51,19 @@ func testMDMWindowsEnrolledDevice(t *testing.T, ds *Datastore) { err = ds.MDMWindowsInsertEnrolledDevice(ctx, enrolledDevice) require.ErrorAs(t, err, &ae) - gotEnrolledDevice, err := ds.MDMWindowsGetEnrolledDevice(ctx, enrolledDevice.MDMDeviceID) + gotEnrolledDevice, err := ds.MDMWindowsGetEnrolledDevice(ctx, enrolledDevice.MDMHardwareID) require.NoError(t, err) require.NotZero(t, gotEnrolledDevice.CreatedAt) require.Equal(t, enrolledDevice.MDMDeviceID, gotEnrolledDevice.MDMDeviceID) require.Equal(t, enrolledDevice.MDMHardwareID, gotEnrolledDevice.MDMHardwareID) - err = ds.MDMWindowsDeleteEnrolledDevice(ctx, enrolledDevice.MDMDeviceID) + err = ds.MDMWindowsDeleteEnrolledDevice(ctx, enrolledDevice.MDMHardwareID) require.NoError(t, err) var nfe fleet.NotFoundError - _, err = ds.MDMWindowsGetEnrolledDevice(ctx, enrolledDevice.MDMDeviceID) + _, err = ds.MDMWindowsGetEnrolledDevice(ctx, enrolledDevice.MDMHardwareID) require.ErrorAs(t, err, &nfe) - err = ds.MDMWindowsDeleteEnrolledDevice(ctx, enrolledDevice.MDMDeviceID) + err = ds.MDMWindowsDeleteEnrolledDevice(ctx, enrolledDevice.MDMHardwareID) require.ErrorAs(t, err, &nfe) } diff --git a/server/fleet/microsoft_mdm.go b/server/fleet/microsoft_mdm.go index 9ff8fe0a8d..c72e37070a 100644 --- a/server/fleet/microsoft_mdm.go +++ b/server/fleet/microsoft_mdm.go @@ -9,6 +9,7 @@ import ( "time" mdm "github.com/fleetdm/fleet/v4/server/mdm/microsoft" + microsoft_mdm "github.com/fleetdm/fleet/v4/server/mdm/microsoft" ) ////////////////////////////////////////////////////////////////////////////////////////////////////////// @@ -44,13 +45,25 @@ type SoapRequest struct { Body BodyRequest `xml:"Body"` } -// GetBinarySecurityToken returns the header BinarySecurityToken if present -func (req *SoapRequest) GetBinarySecurityToken() (string, error) { +// GetHeaderBinarySecurityToken returns the header BinarySecurityToken if present +func (req *SoapRequest) GetHeaderBinarySecurityToken() (*HeaderBinarySecurityToken, error) { if req.Header.Security == nil { - return "", errors.New("header BinarySecurityToken is not present") + return nil, errors.New("binarySecurityToken is not present") } - return req.Header.Security.Security.Content, nil + if len(req.Header.Security.Security.Content) == 0 { + return nil, errors.New("binarySecurityToken is empty") + } + + if req.Header.Security.Security.Encoding != mdm.EnrollEncode { + return nil, errors.New("binarySecurityToken encoding is invalid") + } + + if req.Header.Security.Security.Value != mdm.BinarySecurityDeviceEnroll && req.Header.Security.Security.Value != mdm.BinarySecurityAzureEnroll { + return nil, errors.New("binarySecurityToken type is invalid") + } + + return &req.Header.Security.Security, nil } // GetMessageID returns the message ID from the header @@ -328,16 +341,59 @@ type WsSecurity struct { } // Security token container for encoded security sensitive data -type BinSecurityToken struct { +type HeaderBinarySecurityToken struct { Content string `xml:",chardata"` Value string `xml:"ValueType,attr"` Encoding string `xml:"EncodingType,attr"` } +// Get RequestSecurityToken MDM Message from the body +func (token *HeaderBinarySecurityToken) IsValidToken() error { + if token == nil { + return errors.New("binary security token is not present") + } + + if len(token.Content) == 0 { + return errors.New("binary security token is empty") + } + + if token.Value != microsoft_mdm.BinarySecurityDeviceEnroll && token.Value != microsoft_mdm.BinarySecurityAzureEnroll { + return errors.New("binary security token is invalid") + } + + return nil +} + +// Check if input token is a valid Azure JWT token +func (token *HeaderBinarySecurityToken) IsAzureJWTToken() bool { + if token == nil { + return false + } + + if token.Value == microsoft_mdm.BinarySecurityAzureEnroll { + return true + } + + return false +} + +// Check if input token is a valid Device Enroll token +func (token *HeaderBinarySecurityToken) IsDeviceToken() bool { + if token == nil { + return false + } + + if token.Value == microsoft_mdm.BinarySecurityDeviceEnroll { + return true + } + + return false +} + // TokenSecurity is the security token container for BinSecurityToken type TokenSecurity struct { - MustUnderstand string `xml:"mustUnderstand,attr"` - Security BinSecurityToken `xml:"BinarySecurityToken"` + MustUnderstand string `xml:"mustUnderstand,attr"` + Security HeaderBinarySecurityToken `xml:"BinarySecurityToken"` } // To target endpoint header field @@ -486,10 +542,11 @@ type DiscoverResponse struct { } type DiscoverResult struct { - AuthPolicy string `xml:"AuthPolicy"` - EnrollmentVersion string `xml:"EnrollmentVersion"` - EnrollmentPolicyServiceUrl string `xml:"EnrollmentPolicyServiceUrl"` - EnrollmentServiceUrl string `xml:"EnrollmentServiceUrl"` + AuthPolicy string `xml:"AuthPolicy"` + EnrollmentVersion string `xml:"EnrollmentVersion"` + EnrollmentPolicyServiceUrl string `xml:"EnrollmentPolicyServiceUrl"` + EnrollmentServiceUrl string `xml:"EnrollmentServiceUrl"` + AuthServiceUrl *string `xml:"AuthenticationServiceUrl"` } /////////////////////////////////////////////////////////////// @@ -651,7 +708,8 @@ type WindowsMDMAccessTokenPayload struct { // Type is the enrollment type, such as "programmatic". Type WindowsMDMEnrollmentType `json:"type"` Payload struct { - HostUUID string `json:"host_uuid"` + HostUUID string `json:"host_uuid"` + AuthToken string `json:"auth_token"` } `json:"payload"` } @@ -661,18 +719,23 @@ type WindowsMDMEnrollmentType int const ( WindowsMDMProgrammaticEnrollmentType WindowsMDMEnrollmentType = 1 + WindowsMDMAutomaticEnrollmentType WindowsMDMEnrollmentType = 2 ) func (t *WindowsMDMAccessTokenPayload) IsValidToken() error { // Only BSProgrammaticEnrollment are supported for now - if t.Type != WindowsMDMProgrammaticEnrollmentType { + if t.Type != WindowsMDMProgrammaticEnrollmentType && t.Type != WindowsMDMAutomaticEnrollmentType { return errors.New("invalid binary security payload type") } - if len(t.Payload.HostUUID) == 0 { + if t.Type == WindowsMDMProgrammaticEnrollmentType && len(t.Payload.HostUUID) == 0 { return errors.New("invalid binary security payload content") } + if t.Type == WindowsMDMAutomaticEnrollmentType && len(t.Payload.AuthToken) == 0 { + return errors.New("invalid STS auth token payload content") + } + return nil } diff --git a/server/fleet/service.go b/server/fleet/service.go index 86c6f881c4..c83493ad7e 100644 --- a/server/fleet/service.go +++ b/server/fleet/service.go @@ -760,13 +760,16 @@ type Service interface { // Windows MDM // GetMDMMicrosoftDiscoveryResponse returns a valid DiscoveryResponse message - GetMDMMicrosoftDiscoveryResponse(ctx context.Context) (*DiscoverResponse, error) + GetMDMMicrosoftDiscoveryResponse(ctx context.Context, upnEmail string) (*DiscoverResponse, error) + + // GetMDMMicrosoftSTSAuthResponse returns a valid STS auth page + GetMDMMicrosoftSTSAuthResponse(ctx context.Context, appru string, loginHint string) (string, error) // GetMDMWindowsPolicyResponse returns a valid GetPoliciesResponse message - GetMDMWindowsPolicyResponse(ctx context.Context, authToken string) (*GetPoliciesResponse, error) + GetMDMWindowsPolicyResponse(ctx context.Context, authToken *HeaderBinarySecurityToken) (*GetPoliciesResponse, error) // GetMDMWindowsEnrollResponse returns a valid RequestSecurityTokenResponseCollection message - GetMDMWindowsEnrollResponse(ctx context.Context, secTokenMsg *RequestSecurityToken, authToken string) (*RequestSecurityTokenResponseCollection, error) + GetMDMWindowsEnrollResponse(ctx context.Context, secTokenMsg *RequestSecurityToken, authToken *HeaderBinarySecurityToken) (*RequestSecurityTokenResponseCollection, error) // GetAuthorizedSoapFault authorize the request so SoapFault message can be returned GetAuthorizedSoapFault(ctx context.Context, eType string, origMsg int, errorMsg error) *SoapFault diff --git a/server/mdm/microsoft/microsoft_mdm.go b/server/mdm/microsoft/microsoft_mdm.go index 45ab6b1ae5..590a7c8a1b 100644 --- a/server/mdm/microsoft/microsoft_mdm.go +++ b/server/mdm/microsoft/microsoft_mdm.go @@ -14,6 +14,12 @@ const ( // https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-mde2/2681fd76-1997-4557-8963-cf656ab8d887 MDE2DiscoveryPath = MDMPath + "/discovery" + // AuthPath is the HTTP endpoint path that delivers the Security Token Servicefunctionality. + // The MS-MDE2 protocol is agnostic to the token format and value returned by this endpoint. + // See the section 3.2 on the MS-MDE2 specification for more details: + // https://learn.microsoft.com/en-us/openspecs/windows_protocols/ms-mde2/27ed8c2c-0140-41ce-b2fa-c3d1a793ab4a + MDE2AuthPath = MDMPath + "/auth" + // MDE2PolicyPath is the HTTP endpoint path that delivers the X.509 Certificate Enrollment Policy (MS-XCEP) functionality. // This is the endpoint that process the GetPolicies and GetPoliciesResponse messages // See the section 3.3 on the MS-MDE2 specification for more details on this endpoint requirements: @@ -42,27 +48,29 @@ const ( MSManageEntryPoint = "/ManagementServer/MDM.svc" ) -// XML Namespaces used by the Microsoft Device Enrollment v2 protocol (MS-MDE2) +// XML Namespaces and type URLs used by the Microsoft Device Enrollment v2 protocol (MS-MDE2) const ( - DiscoverNS = "http://schemas.microsoft.com/windows/management/2012/01/enrollment" - PolicyNS = "http://schemas.microsoft.com/windows/pki/2009/01/enrollmentpolicy" - EnrollWSTrust = "http://docs.oasis-open.org/ws-sx/ws-trust/200512" - EnrollSecExt = "http://docs.oasis-open.org/wss/2004/01/oasis-200401-wss-wssecurity-secext-1.0.xsd" - EnrollTType = "http://schemas.microsoft.com/5.0.0.0/ConfigurationManager/Enrollment/DeviceEnrollmentToken" - EnrollPDoc = "http://schemas.microsoft.com/5.0.0.0/ConfigurationManager/Enrollment/DeviceEnrollmentProvisionDoc" - EnrollEncode = "http://docs.oasis-open.org/wss/2004/01/oasis-200401-wss-wssecurity-secext-1.0.xsd#base64binary" - EnrollReq = "http://schemas.microsoft.com/windows/pki/2009/01/enrollment" - EnrollNSS = "http://www.w3.org/2003/05/soap-envelope" - EnrollNSA = "http://www.w3.org/2005/08/addressing" - EnrollXSI = "http://www.w3.org/2001/XMLSchema-instance" - EnrollXSD = "http://www.w3.org/2001/XMLSchema" - EnrollXSU = "http://docs.oasis-open.org/wss/2004/01/oasis-200401-wss-wssecurity-utility-1.0.xsd" - ActionNsDiag = "http://schemas.microsoft.com/2004/09/ServiceModel/Diagnostics" - ActionNsDiscovery = "http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/DiscoverResponse" - ActionNsPolicy = "http://schemas.microsoft.com/windows/pki/2009/01/enrollmentpolicy/IPolicy/GetPoliciesResponse" - ActionNsEnroll = EnrollReq + "/RSTRC/wstep" - EnrollReqTypePKCS10 = EnrollReq + "#PKCS10" - EnrollReqTypePKCS7 = EnrollReq + "#PKCS7" + DiscoverNS = "http://schemas.microsoft.com/windows/management/2012/01/enrollment" + PolicyNS = "http://schemas.microsoft.com/windows/pki/2009/01/enrollmentpolicy" + EnrollWSTrust = "http://docs.oasis-open.org/ws-sx/ws-trust/200512" + EnrollSecExt = "http://docs.oasis-open.org/wss/2004/01/oasis-200401-wss-wssecurity-secext-1.0.xsd" + EnrollTType = "http://schemas.microsoft.com/5.0.0.0/ConfigurationManager/Enrollment/DeviceEnrollmentToken" + EnrollPDoc = "http://schemas.microsoft.com/5.0.0.0/ConfigurationManager/Enrollment/DeviceEnrollmentProvisionDoc" + EnrollEncode = "http://docs.oasis-open.org/wss/2004/01/oasis-200401-wss-wssecurity-secext-1.0.xsd#base64binary" + EnrollReq = "http://schemas.microsoft.com/windows/pki/2009/01/enrollment" + EnrollNSS = "http://www.w3.org/2003/05/soap-envelope" + EnrollNSA = "http://www.w3.org/2005/08/addressing" + EnrollXSI = "http://www.w3.org/2001/XMLSchema-instance" + EnrollXSD = "http://www.w3.org/2001/XMLSchema" + EnrollXSU = "http://docs.oasis-open.org/wss/2004/01/oasis-200401-wss-wssecurity-utility-1.0.xsd" + ActionNsDiag = "http://schemas.microsoft.com/2004/09/ServiceModel/Diagnostics" + ActionNsDiscovery = "http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/DiscoverResponse" + ActionNsPolicy = "http://schemas.microsoft.com/windows/pki/2009/01/enrollmentpolicy/IPolicy/GetPoliciesResponse" + ActionNsEnroll = EnrollReq + "/RSTRC/wstep" + EnrollReqTypePKCS10 = EnrollReq + "#PKCS10" + EnrollReqTypePKCS7 = EnrollReq + "#PKCS7" + BinarySecurityDeviceEnroll = "http://schemas.microsoft.com/5.0.0.0/ConfigurationManager/Enrollment/DeviceEnrollmentUserToken" + BinarySecurityAzureEnroll = "urn:ietf:params:oauth:token-type:jwt" ) // Soap Error constants @@ -214,6 +222,12 @@ const ( ReqSecTokenContextItemApplicationVersion = "ApplicationVersion" ReqSecTokenContextItemNotInOobe = "NotInOobe" ReqSecTokenContextItemRequestVersion = "RequestVersion" + + // APPRU query param expected by STS Auth endpoint + STSAuthAppRu = "appru" + + // Login related query param expected by STS Auth endpoint + STSLoginHint = "login_hint" ) func ResolveWindowsMDMDiscovery(serverURL string) (string, error) { @@ -228,6 +242,10 @@ func ResolveWindowsMDMEnroll(serverURL string) (string, error) { return commonmdm.ResolveURL(serverURL, MDE2EnrollPath, false) } +func ResolveWindowsMDMAuth(serverURL string) (string, error) { + return commonmdm.ResolveURL(serverURL, MDE2AuthPath, false) +} + func ResolveWindowsMDMManagement(serverURL string) (string, error) { return commonmdm.ResolveURL(serverURL, MDE2ManagementPath, false) } diff --git a/server/mdm/microsoft/wstep.go b/server/mdm/microsoft/wstep.go index 9171218bbc..54f8fb009d 100644 --- a/server/mdm/microsoft/wstep.go +++ b/server/mdm/microsoft/wstep.go @@ -1,6 +1,7 @@ package microsoft_mdm import ( + "bytes" "context" "crypto/rand" "crypto/rsa" @@ -17,6 +18,7 @@ import ( "time" "github.com/fleetdm/fleet/v4/server" + "github.com/golang-jwt/jwt/v4" "github.com/micromdm/nanomdm/cryptoutil" "go.mozilla.org/pkcs7" ) @@ -35,6 +37,12 @@ type CertManager interface { // IdentityCert returns the identity certificate of the depot. IdentityCert() x509.Certificate + // NewSTSAuthToken returns an STS auth token for the given UPN claim. + NewSTSAuthToken(upn string) (string, error) + + // GetSTSAuthTokenUPNClaim validates the given token and returns the UPN claim + GetSTSAuthTokenUPNClaim(token string) (string, error) + // TODO: implement other methods as needed: // - verify certificate-device association // - certificate lifecycle management (e.g., renewal, revocation) @@ -48,6 +56,18 @@ type CertStore interface { WSTEPAssociateCertHash(ctx context.Context, deviceUUID string, hash string) error } +type STSClaims struct { + UPN string `json:"upn"` + jwt.RegisteredClaims +} + +type AzureData struct { + UPN string + TenantID string + UniqueName string + SCP string +} + type manager struct { store CertStore @@ -144,6 +164,134 @@ func (m *manager) SignClientCSR(ctx context.Context, subject string, clientCSR * return rawSignedDER, CertFingerprintHexStr(signedCert), nil } +// NewSTSAuthToken returns an STS auth token for the given UPN claim. +func (m *manager) NewSTSAuthToken(upn string) (string, error) { + if m == nil { + return "", errors.New("windows mdm identity keypair was not configured") + } + + if m.identityCert == nil || m.identityPrivateKey == nil { + return "", errors.New("invalid identity certificate or private key") + } + + if len(upn) == 0 { + return "", errors.New("invalid upn field") + } + + // Create claims with upn field populated + claims := STSClaims{ + upn, + jwt.RegisteredClaims{ + ExpiresAt: jwt.NewNumericDate(time.Now().Add(10 * time.Minute)), + IssuedAt: jwt.NewNumericDate(time.Now()), + NotBefore: jwt.NewNumericDate(time.Now()), + Subject: "STSAuthToken", + }, + } + + // Create a new token with the claims and sign it with the private key + token := jwt.NewWithClaims(jwt.GetSigningMethod("RS256"), claims) + signedToken, err := token.SignedString(m.identityPrivateKey) + if err != nil { + return "", fmt.Errorf("failed to sign STS token: %w", err) + } + + return signedToken, nil +} + +// GetSTSAuthToken validates the given token and returns the UPN claim +func (m *manager) GetSTSAuthTokenUPNClaim(tokenStr string) (string, error) { + if m == nil { + return "", errors.New("windows mdm identity keypair was not configured") + } + + if m.identityCert == nil || m.identityPrivateKey == nil { + return "", errors.New("invalid identity certificate or private key") + } + + if len(tokenStr) == 0 { + return "", errors.New("invalid STS token") + } + + // Since we used the private key to sign the tokens, we use the public counterpart to verify the signature + token, err := jwt.ParseWithClaims(tokenStr, &STSClaims{}, func(token *jwt.Token) (interface{}, error) { + return m.identityCert.PublicKey, nil + }) + if err != nil { + return "", fmt.Errorf("there was an error parsing the STS token claims: %w", err) + } + + if claims, ok := token.Claims.(*STSClaims); ok && token.Valid { + if len(claims.UPN) == 0 { + return "", errors.New("issue with UPN token claim") + } + + return claims.UPN, nil + } + + return "", errors.New("issue with STS token validation") +} + +// GetAzureAuthTokenClaims validates the given Azure AD token and returns +// UPN, TenantID, UniqueName, DeviceID +func GetAzureAuthTokenClaims(tokenStr string) (AzureData, error) { + if len(tokenStr) == 0 { + return AzureData{}, errors.New("invalid STS token") + } + + // Decode base64 token + tokenBytes, err := base64.StdEncoding.DecodeString(tokenStr) + if err != nil { + return AzureData{}, errors.New("invalid Azure JWT token") + } + + // Validate token format (header.payload.signature) + parts := bytes.Split(tokenBytes, []byte(".")) + if len(parts) != 3 { + return AzureData{}, errors.New("invalid Azure JWT format") + } + + // Parse JWT token + token, _, err := new(jwt.Parser).ParseUnverified(string(tokenBytes), jwt.MapClaims{}) + if err != nil { + return AzureData{}, errors.New("parse error Azure JWT content") + } + + // Parse JWT token + claims := token.Claims.(jwt.MapClaims) + + // Get UPN claim + upnClaim, ok := claims["upn"].(string) + if !ok || len(upnClaim) == 0 { + return AzureData{}, errors.New("invalid UPN claim") + } + + // Get TenantID claim + tenantIDClaim, ok := claims["tid"].(string) + if !ok || len(tenantIDClaim) == 0 { + return AzureData{}, errors.New("invalid TenantID claim") + } + + // Get UniqueName claim + uniqueNameClaim, ok := claims["unique_name"].(string) + if !ok { + return AzureData{}, errors.New("invalid UniqueName claim") + } + + // Get SCP claim + azureSCPClaim, ok := claims["scp"].(string) + if !ok || azureSCPClaim != "mdm_delegation" { + return AzureData{}, errors.New("invalid SCP claim") + } + + return AzureData{ + UPN: upnClaim, + TenantID: tenantIDClaim, + UniqueName: uniqueNameClaim, + SCP: azureSCPClaim, + }, nil +} + func populateClientCert(sn *big.Int, subject string, issuerCert *x509.Certificate, csr *x509.CertificateRequest) (*x509.Certificate, error) { certRenewalPeriodInSecsInt, err := strconv.Atoi(PolicyCertRenewalPeriodInSecs) if err != nil { diff --git a/server/mdm/microsoft/wstep_test.go b/server/mdm/microsoft/wstep_test.go index 66c202a5b3..298b8a183b 100644 --- a/server/mdm/microsoft/wstep_test.go +++ b/server/mdm/microsoft/wstep_test.go @@ -75,12 +75,28 @@ func TestNewCertManager(t *testing.T) { require.Equal(t, wantIdentityFingerprint, m.identityFingerprint) } -func TestSignClientCSR(t *testing.T) { - // TODO -} +func TestSTSTokenSigningAndVerification(t *testing.T) { + var store CertStore -func TestGetClientCSR(t *testing.T) { - // TODO + cm, err := NewCertManager(store, testCert, testKey) + require.NoError(t, err) + require.NotNil(t, cm) + + // Get a New STS Auth token + upnEmail := "test@email.com" + stsToken, err := cm.NewSTSAuthToken(upnEmail) + require.NoError(t, err) + require.NotEmpty(t, stsToken) + + // Verify the STS Auth token + upnToken, err := cm.GetSTSAuthTokenUPNClaim(stsToken) + require.NoError(t, err) + require.NotEmpty(t, upnToken) + require.Equal(t, upnEmail, upnToken) + + // New invalid STS Auth token + _, err = cm.NewSTSAuthToken("") + require.ErrorContains(t, err, "invalid upn field") } func TestCertFingerprintHexStr(t *testing.T) { diff --git a/server/service/endpoint_utils.go b/server/service/endpoint_utils.go index cfe19c22b9..db63c49953 100644 --- a/server/service/endpoint_utils.go +++ b/server/service/endpoint_utils.go @@ -10,6 +10,7 @@ import ( "io" "net" "net/http" + "net/url" "reflect" "strconv" "strings" @@ -83,10 +84,9 @@ type requestDecoder interface { } // A value that implements bodyDecoder takes control of decoding the request -// body. Other fields such as url and query parameters are decoded prior to -// calling DecodeBody with the request's body as an io.Reader. +// body. type bodyDecoder interface { - DecodeBody(ctx context.Context, r io.Reader) error + DecodeBody(ctx context.Context, r io.Reader, u url.Values) error } // makeDecoder creates a decoder for the type for the struct passed on. If the @@ -304,7 +304,7 @@ func makeDecoder(iface interface{}) kithttp.DecodeRequestFunc { if isBodyDecoder { bd := v.Interface().(bodyDecoder) - if err := bd.DecodeBody(ctx, body); err != nil { + if err := bd.DecodeBody(ctx, body, r.URL.Query()); err != nil { return nil, err } } diff --git a/server/service/handler.go b/server/service/handler.go index a0b6c1c1fa..a076b2649e 100644 --- a/server/service/handler.go +++ b/server/service/handler.go @@ -597,10 +597,13 @@ func attachFleetAPIRoutes(r *mux.Router, svc fleet.Service, config config.FleetC // These endpoint are used by Microsoft devices during MDM device enrollment phase neWindowsMDM := ne.WithCustomMiddleware(mdmConfiguredMiddleware.VerifyWindowsMDM()) - // Microsoft MS-MDE Endpoints - // This endpoint is unauthenticated and is used by Microsoft devices to discover the MDM server + // Microsoft MS-MDE2 Endpoints + // This endpoint is unauthenticated and is used by Microsoft devices to discover the MDM server endpoints neWindowsMDM.POST(microsoft_mdm.MDE2DiscoveryPath, mdmMicrosoftDiscoveryEndpoint, SoapRequestContainer{}) + // This endpoint is unauthenticated and is used by Microsoft devices to retrieve the opaque STS auth token + neWindowsMDM.GET(microsoft_mdm.MDE2AuthPath, mdmMicrosoftAuthEndpoint, SoapRequestContainer{}) + // This endpoint is authenticated using the BinarySecurityToken header field neWindowsMDM.POST(microsoft_mdm.MDE2PolicyPath, mdmMicrosoftPolicyEndpoint, SoapRequestContainer{}) diff --git a/server/service/integration_mdm_test.go b/server/service/integration_mdm_test.go index 8cfc77feec..b15b7485dd 100644 --- a/server/service/integration_mdm_test.go +++ b/server/service/integration_mdm_test.go @@ -5574,331 +5574,6 @@ func (s *integrationMDMTestSuite) TestAppConfigWindowsMDM() { require.Empty(t, resp.Notifications.WindowsMDMDiscoveryEndpoint) } -func (s *integrationMDMTestSuite) TestValidDiscoveryRequest() { - t := s.T() - - // Preparing the Discovery Request message - requestBytes := []byte(` - - - http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/Discover - urn:uuid:148132ec-a575-4322-b01b-6172a9cf8478 - - http://www.w3.org/2005/08/addressing/anonymous - - https://mdmwindows.com:443/EnrollmentServer/Discovery.svc - - - - - demo@mdmwindows.com - 5.0 - CIMClient_Windows - 6.2.9200.2965 - 48 - - OnPremise - Federated - - - - - `) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2DiscoveryPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if SOAP response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid DiscoveryResponse message - resSoapMsg := string(resBytes) - require.True(t, s.isXMLTagPresent("DiscoverResult", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("AuthPolicy", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("EnrollmentVersion", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("EnrollmentPolicyServiceUrl", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("EnrollmentServiceUrl", resSoapMsg)) -} - -func (s *integrationMDMTestSuite) TestInvalidDiscoveryRequest() { - t := s.T() - - // Preparing the Discovery Request message - requestBytes := []byte(` - - - http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/Discover - - http://www.w3.org/2005/08/addressing/anonymous - - https://mdmwindows.com:443/EnrollmentServer/Discovery.svc - - - - - demo@mdmwindows.com - 5.0 - CIMClient_Windows - 6.2.9200.2965 - 48 - - OnPremise - Federated - - - - - `) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2DiscoveryPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid SoapFault message - resSoapMsg := string(resBytes) - - require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) - require.True(t, s.checkIfXMLTagContains("s:text", "invalid SOAP header: Header.MessageID", resSoapMsg)) -} - -func (s *integrationMDMTestSuite) TestValidGetPoliciesRequest() { - t := s.T() - - // create a new Host to get the UUID on the DB - windowsHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ - ID: 1, - OsqueryHostID: ptr.String("Desktop-ABCQWE"), - NodeKey: ptr.String("Desktop-ABCQWE"), - UUID: uuid.New().String(), - Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), - Platform: "windows", - }) - require.NoError(t, err) - - // Preparing the GetPolicies Request message - encodedBinToken, err := GetEncodedBinarySecurityToken(1, windowsHost.UUID) - require.NoError(t, err) - - requestBytes, err := s.newGetPoliciesMsg(encodedBinToken) - require.NoError(t, err) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if SOAP response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid GetPoliciesResponse message - resSoapMsg := string(resBytes) - require.True(t, s.isXMLTagPresent("GetPoliciesResponse", resSoapMsg)) - require.True(t, s.isXMLTagPresent("policyOIDReference", resSoapMsg)) - require.True(t, s.isXMLTagPresent("oIDReferenceID", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("validityPeriodSeconds", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("renewalPeriodSeconds", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("minimalKeyLength", resSoapMsg)) -} - -func (s *integrationMDMTestSuite) TestGetPoliciesRequestWithInvalidUUID() { - t := s.T() - - // create a new Host to get the UUID on the DB - _, err := s.ds.NewHost(context.Background(), &fleet.Host{ - ID: 1, - OsqueryHostID: ptr.String("Desktop-ABCQWE"), - NodeKey: ptr.String("Desktop-ABCQWE"), - UUID: uuid.New().String(), - Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), - Platform: "windows", - }) - require.NoError(t, err) - - // Preparing the GetPolicies Request message - encodedBinToken, err := GetEncodedBinarySecurityToken(1, "not_exists") - require.NoError(t, err) - - requestBytes, err := s.newGetPoliciesMsg(encodedBinToken) - require.NoError(t, err) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if SOAP response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid SoapFault message - resSoapMsg := string(resBytes) - require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) - require.True(t, s.checkIfXMLTagContains("s:text", "binarySecurityTokenValidation: host data cannot be found", resSoapMsg)) -} - -func (s *integrationMDMTestSuite) TestGetPoliciesRequestWithNotElegibleHost() { - t := s.T() - - // create a new Host to get the UUID on the DB - linuxHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ - ID: 1, - OsqueryHostID: ptr.String("Ubuntu01"), - NodeKey: ptr.String("Ubuntu01"), - UUID: uuid.New().String(), - Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), - Platform: "linux", - }) - require.NoError(t, err) - - // Preparing the GetPolicies Request message - encodedBinToken, err := GetEncodedBinarySecurityToken(1, linuxHost.UUID) - require.NoError(t, err) - - requestBytes, err := s.newGetPoliciesMsg(encodedBinToken) - require.NoError(t, err) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if SOAP response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid SoapFault message - resSoapMsg := string(resBytes) - require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) - require.True(t, s.checkIfXMLTagContains("s:text", "host is not elegible for Windows MDM enrollment", resSoapMsg)) -} - -func (s *integrationMDMTestSuite) TestValidRequestSecurityTokenRequest() { - t := s.T() - - // create a new Host to get the UUID on the DB - windowsHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ - ID: 1, - OsqueryHostID: ptr.String("Desktop-ABCQWE"), - NodeKey: ptr.String("Desktop-ABCQWE"), - UUID: uuid.New().String(), - Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), - Platform: "windows", - }) - require.NoError(t, err) - - // Delete the host from the list of MDM enrolled devices if present - _ = s.ds.MDMWindowsDeleteEnrolledDevice(context.Background(), windowsHost.UUID) - - // Preparing the RequestSecurityToken Request message - encodedBinToken, err := GetEncodedBinarySecurityToken(1, windowsHost.UUID) - require.NoError(t, err) - - requestBytes, err := s.newSecurityTokenMsg(encodedBinToken, true) - require.NoError(t, err) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2EnrollPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if SOAP response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid RequestSecurityTokenResponseCollection message - resSoapMsg := string(resBytes) - require.True(t, s.isXMLTagPresent("RequestSecurityTokenResponseCollection", resSoapMsg)) - require.True(t, s.isXMLTagPresent("DispositionMessage", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("TokenType", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("RequestID", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("BinarySecurityToken", resSoapMsg)) - - // Checking if an activity was created for the enrollment - s.lastActivityOfTypeMatches( - fleet.ActivityTypeMDMEnrolled{}.ActivityName(), - `{ - "mdm_platform": "microsoft", - "host_serial": "", - "installed_from_dep": false, - "host_display_name": "DESKTOP-0C89RC0" - }`, - 0) -} - -func (s *integrationMDMTestSuite) TestInvalidRequestSecurityTokenRequestWithMissingAdditionalContext() { - t := s.T() - - // create a new Host to get the UUID on the DB - windowsHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ - ID: 1, - OsqueryHostID: ptr.String("Desktop-ABCQWE"), - NodeKey: ptr.String("Desktop-ABCQWE"), - UUID: uuid.New().String(), - Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), - Platform: "windows", - }) - require.NoError(t, err) - - // Preparing the RequestSecurityToken Request message - encodedBinToken, err := GetEncodedBinarySecurityToken(1, windowsHost.UUID) - require.NoError(t, err) - - requestBytes, err := s.newSecurityTokenMsg(encodedBinToken, false) - require.NoError(t, err) - - resp := s.DoRaw("POST", microsoft_mdm.MDE2EnrollPath, requestBytes, http.StatusOK) - - resBytes, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) - - // Checking if SOAP response can be unmarshalled to an golang type - var xmlType interface{} - err = xml.Unmarshal(resBytes, &xmlType) - require.NoError(t, err) - - // Checking if SOAP response contains a valid SoapFault message - resSoapMsg := string(resBytes) - require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) - require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) - require.True(t, s.checkIfXMLTagContains("s:text", "ContextItem item DeviceType is not present", resSoapMsg)) -} - func (s *integrationMDMTestSuite) TestOrbitConfigNudgeSettings() { t := s.T() @@ -6009,6 +5684,492 @@ func (s *integrationMDMTestSuite) TestOrbitConfigNudgeSettings() { require.Equal(t, wantCfg.OSVersionRequirements[0].RequiredInstallationDate.String(), "2022-01-04 04:00:00 +0000 UTC") } +func (s *integrationMDMTestSuite) TestValidDiscoveryRequest() { + t := s.T() + + // Preparing the Discovery Request message + requestBytes := []byte(` + + + http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/Discover + urn:uuid:148132ec-a575-4322-b01b-6172a9cf8478 + + http://www.w3.org/2005/08/addressing/anonymous + + https://mdmwindows.com:443/EnrollmentServer/Discovery.svc + + + + + demo@mdmwindows.com + 5.0 + CIMClient_Windows + 6.2.9200.2965 + 48 + + OnPremise + Federated + + + + + `) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2DiscoveryPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid DiscoveryResponse message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("DiscoverResult", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("AuthPolicy", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("EnrollmentVersion", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("EnrollmentPolicyServiceUrl", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("EnrollmentServiceUrl", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestInvalidDiscoveryRequest() { + t := s.T() + + // Preparing the Discovery Request message + requestBytes := []byte(` + + + http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/Discover + + http://www.w3.org/2005/08/addressing/anonymous + + https://mdmwindows.com:443/EnrollmentServer/Discovery.svc + + + + + demo@mdmwindows.com + 5.0 + CIMClient_Windows + 6.2.9200.2965 + 48 + + OnPremise + Federated + + + + + `) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2DiscoveryPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid SoapFault message + resSoapMsg := string(resBytes) + + require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) + require.True(t, s.checkIfXMLTagContains("s:text", "invalid SOAP header: Header.MessageID", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestNoEmailDiscoveryRequest() { + t := s.T() + + // Preparing the Discovery Request message + requestBytes := []byte(` + + + http://schemas.microsoft.com/windows/management/2012/01/enrollment/IDiscoveryService/Discover + urn:uuid:148132ec-a575-4322-b01b-6172a9cf8478 + + http://www.w3.org/2005/08/addressing/anonymous + + https://mdmwindows.com:443/EnrollmentServer/Discovery.svc + + + + + + 5.0 + CIMClient_Windows + 6.2.9200.2965 + 48 + + OnPremise + Federated + + + + + `) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2DiscoveryPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid DiscoveryResponse message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("DiscoverResult", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("AuthPolicy", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("EnrollmentVersion", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("EnrollmentPolicyServiceUrl", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("EnrollmentServiceUrl", resSoapMsg)) + require.True(t, !s.isXMLTagContentPresent("AuthenticationServiceUrl", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestValidGetPoliciesRequestWithDeviceToken() { + t := s.T() + + // create a new Host to get the UUID on the DB + windowsHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ + ID: 1, + OsqueryHostID: ptr.String("Desktop-ABCQWE"), + NodeKey: ptr.String("Desktop-ABCQWE"), + UUID: uuid.New().String(), + Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), + Platform: "windows", + }) + require.NoError(t, err) + + // Preparing the GetPolicies Request message + encodedBinToken, err := GetEncodedBinarySecurityToken(fleet.WindowsMDMProgrammaticEnrollmentType, windowsHost.UUID) + require.NoError(t, err) + + requestBytes, err := s.newGetPoliciesMsg(true, encodedBinToken) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid GetPoliciesResponse message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("GetPoliciesResponse", resSoapMsg)) + require.True(t, s.isXMLTagPresent("policyOIDReference", resSoapMsg)) + require.True(t, s.isXMLTagPresent("oIDReferenceID", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("validityPeriodSeconds", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("renewalPeriodSeconds", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("minimalKeyLength", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestValidGetPoliciesRequestWithAzureToken() { + t := s.T() + + // Preparing the GetPolicies Request message with Azure JWT token + azureADTok := "ZXlKMGVYQWlPaUpLVjFRaUxDSmhiR2NpT2lKU1V6STFOaUlzSW5nMWRDSTZJaTFMU1ROUk9XNU9VamRpVW05bWVHMWxXbTlZY1dKSVdrZGxkeUlzSW10cFpDSTZJaTFMU1ROUk9XNU9VamRpVW05bWVHMWxXbTlZY1dKSVdrZGxkeUo5LmV5SmhkV1FpT2lKb2RIUndjem92TDIxaGNtTnZjMnhoWW5NdWIzSm5MeUlzSW1semN5STZJbWgwZEhCek9pOHZjM1J6TG5kcGJtUnZkM011Ym1WMEwyWmhaVFZqTkdZekxXWXpNVGd0TkRRNE15MWlZelptTFRjMU9UVTFaalJoTUdFM01pOGlMQ0pwWVhRaU9qRTJPRGt4TnpBNE5UZ3NJbTVpWmlJNk1UWTRPVEUzTURnMU9Dd2laWGh3SWpveE5qZzVNVGMxTmpZeExDSmhZM0lpT2lJeElpd2lZV2x2SWpvaVFWUlJRWGt2T0ZSQlFVRkJOV2gwUTNFMGRERjNjbHBwUTIxQmVEQlpWaTloZGpGTVMwRkRPRXM1Vm10SGVtNUdXVGxzTUZoYWVrZHVha2N6VVRaMWVIUldNR3QxT1hCeFJXdFRZeUlzSW1GdGNpSTZXeUp3ZDJRaUxDSnljMkVpWFN3aVlYQndhV1FpT2lJeU9XUTVaV1E1T0MxaE5EWTVMVFExTXpZdFlXUmxNaTFtT1RneFltTXhaRFl3TldVaUxDSmhjSEJwWkdGamNpSTZJakFpTENKa1pYWnBZMlZwWkNJNkltRXhNMlkzWVdVd0xURXpPR0V0TkdKaU1pMDVNalF5TFRka09USXlaVGRqTkdGak15SXNJbWx3WVdSa2NpSTZJakU0Tmk0eE1pNHhPRGN1TWpZaUxDSnVZVzFsSWpvaVZHVnpkRTFoY21OdmMweGhZbk1pTENKdmFXUWlPaUpsTTJNMU5XVmtZeTFqTXpRNExUUTBNVFl0T0dZd05TMHlOVFJtWmpNd05qVmpOV1VpTENKd2QyUmZkWEpzSWpvaWFIUjBjSE02THk5d2IzSjBZV3d1YldsamNtOXpiMlowYjI1c2FXNWxMbU52YlM5RGFHRnVaMlZRWVhOemQyOXlaQzVoYzNCNElpd2ljbWdpT2lJd0xrRldTVUU0T0ZSc0xXaHFlbWN3VXpoaU0xZFdXREJ2UzJOdFZGRXpTbHB1ZUUxa1QzQTNUbVZVVm5OV2FYVkhOa0ZRYnk0aUxDSnpZM0FpT2lKdFpHMWZaR1ZzWldkaGRHbHZiaUlzSW5OMVlpSTZJa1pTUTJ4RldURk9ObXR2ZEdWblMzcFplV0pFTjJkdFdGbGxhVTVIUkZrd05FSjJOV3R6ZDJGeGJVRWlMQ0owYVdRaU9pSm1ZV1UxWXpSbU15MW1NekU0TFRRME9ETXRZbU0yWmkwM05UazFOV1kwWVRCaE56SWlMQ0oxYm1seGRXVmZibUZ0WlNJNkluUmxjM1JBYldGeVkyOXpiR0ZpY3k1dmNtY2lMQ0oxY0c0aU9pSjBaWE4wUUcxaGNtTnZjMnhoWW5NdWIzSm5JaXdpZFhScElqb2lNVGg2WkVWSU5UZFRSWFZyYWpseGJqRm9aMlJCUVNJc0luWmxjaUk2SWpFdU1DSjkuVG1FUlRsZktBdWo5bTVvQUc2UTBRblV4VEFEaTNFamtlNHZ3VXo3UTdqUUFVZVZGZzl1U0pzUXNjU2hFTXVxUmQzN1R2VlpQanljdEVoRFgwLVpQcEVVYUlSempuRVEyTWxvc21SZURYZzhrYkhNZVliWi1jb0ZucDEyQkVpQnpJWFBGZnBpaU1GRnNZZ0hSSF9tSWxwYlBlRzJuQ2p0LTZSOHgzYVA5QS1tM0J3eV91dnV0WDFNVEVZRmFsekhGa04wNWkzbjZRcjhURnlJQ1ZUYW5OanlkMjBBZFRMbHJpTVk0RVBmZzRaLThVVTctZkcteElycWVPUmVWTnYwOUFHV192MDd6UkVaNmgxVk9tNl9nelRGcElVVURuZFdabnFLTHlySDlkdkF3WnFFSG1HUmlTNElNWnRFdDJNTkVZSnhDWHhlSi1VbWZJdV9tUVhKMW9R" + requestBytes, err := s.newGetPoliciesMsg(false, azureADTok) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid GetPoliciesResponse message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("GetPoliciesResponse", resSoapMsg)) + require.True(t, s.isXMLTagPresent("policyOIDReference", resSoapMsg)) + require.True(t, s.isXMLTagPresent("oIDReferenceID", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("validityPeriodSeconds", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("renewalPeriodSeconds", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("minimalKeyLength", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestGetPoliciesRequestWithInvalidUUID() { + t := s.T() + + // create a new Host to get the UUID on the DB + _, err := s.ds.NewHost(context.Background(), &fleet.Host{ + ID: 1, + OsqueryHostID: ptr.String("Desktop-ABCQWE"), + NodeKey: ptr.String("Desktop-ABCQWE"), + UUID: uuid.New().String(), + Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), + Platform: "windows", + }) + require.NoError(t, err) + + // Preparing the GetPolicies Request message + encodedBinToken, err := GetEncodedBinarySecurityToken(fleet.WindowsMDMProgrammaticEnrollmentType, "not_exists") + require.NoError(t, err) + + requestBytes, err := s.newGetPoliciesMsg(true, encodedBinToken) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid SoapFault message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) + require.True(t, s.checkIfXMLTagContains("s:text", "host data cannot be found", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestGetPoliciesRequestWithNotElegibleHost() { + t := s.T() + + // create a new Host to get the UUID on the DB + linuxHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ + ID: 1, + OsqueryHostID: ptr.String("Ubuntu01"), + NodeKey: ptr.String("Ubuntu01"), + UUID: uuid.New().String(), + Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), + Platform: "linux", + }) + require.NoError(t, err) + + // Preparing the GetPolicies Request message + encodedBinToken, err := GetEncodedBinarySecurityToken(fleet.WindowsMDMProgrammaticEnrollmentType, linuxHost.UUID) + require.NoError(t, err) + + requestBytes, err := s.newGetPoliciesMsg(true, encodedBinToken) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2PolicyPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid SoapFault message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) + require.True(t, s.checkIfXMLTagContains("s:text", "host is not elegible for Windows MDM enrollment", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestValidRequestSecurityTokenRequestWithDeviceToken() { + t := s.T() + + // create a new Host to get the UUID on the DB + windowsHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ + ID: 1, + OsqueryHostID: ptr.String("Desktop-ABCQWE"), + NodeKey: ptr.String("Desktop-ABCQWE"), + UUID: uuid.New().String(), + Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), + Platform: "windows", + }) + require.NoError(t, err) + + // Delete the host from the list of MDM enrolled devices if present + _ = s.ds.MDMWindowsDeleteEnrolledDevice(context.Background(), windowsHost.UUID) + + // Preparing the RequestSecurityToken Request message + encodedBinToken, err := GetEncodedBinarySecurityToken(fleet.WindowsMDMProgrammaticEnrollmentType, windowsHost.UUID) + require.NoError(t, err) + + requestBytes, err := s.newSecurityTokenMsg(encodedBinToken, true, false) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2EnrollPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid RequestSecurityTokenResponseCollection message + resSoapMsg := string(resBytes) + + require.True(t, s.isXMLTagPresent("RequestSecurityTokenResponseCollection", resSoapMsg)) + require.True(t, s.isXMLTagPresent("DispositionMessage", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("TokenType", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("RequestID", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("BinarySecurityToken", resSoapMsg)) + + // Checking if an activity was created for the enrollment + s.lastActivityOfTypeMatches( + fleet.ActivityTypeMDMEnrolled{}.ActivityName(), + `{ + "mdm_platform": "microsoft", + "host_serial": "", + "installed_from_dep": false, + "host_display_name": "DESKTOP-0C89RC0" + }`, + 0) +} + +func (s *integrationMDMTestSuite) TestValidRequestSecurityTokenRequestWithAzureToken() { + t := s.T() + + // Preparing the SecurityToken Request message with Azure JWT token + azureADTok := "ZXlKMGVYQWlPaUpLVjFRaUxDSmhiR2NpT2lKU1V6STFOaUlzSW5nMWRDSTZJaTFMU1ROUk9XNU9VamRpVW05bWVHMWxXbTlZY1dKSVdrZGxkeUlzSW10cFpDSTZJaTFMU1ROUk9XNU9VamRpVW05bWVHMWxXbTlZY1dKSVdrZGxkeUo5LmV5SmhkV1FpT2lKb2RIUndjem92TDIxaGNtTnZjMnhoWW5NdWIzSm5MeUlzSW1semN5STZJbWgwZEhCek9pOHZjM1J6TG5kcGJtUnZkM011Ym1WMEwyWmhaVFZqTkdZekxXWXpNVGd0TkRRNE15MWlZelptTFRjMU9UVTFaalJoTUdFM01pOGlMQ0pwWVhRaU9qRTJPRGt4TnpBNE5UZ3NJbTVpWmlJNk1UWTRPVEUzTURnMU9Dd2laWGh3SWpveE5qZzVNVGMxTmpZeExDSmhZM0lpT2lJeElpd2lZV2x2SWpvaVFWUlJRWGt2T0ZSQlFVRkJOV2gwUTNFMGRERjNjbHBwUTIxQmVEQlpWaTloZGpGTVMwRkRPRXM1Vm10SGVtNUdXVGxzTUZoYWVrZHVha2N6VVRaMWVIUldNR3QxT1hCeFJXdFRZeUlzSW1GdGNpSTZXeUp3ZDJRaUxDSnljMkVpWFN3aVlYQndhV1FpT2lJeU9XUTVaV1E1T0MxaE5EWTVMVFExTXpZdFlXUmxNaTFtT1RneFltTXhaRFl3TldVaUxDSmhjSEJwWkdGamNpSTZJakFpTENKa1pYWnBZMlZwWkNJNkltRXhNMlkzWVdVd0xURXpPR0V0TkdKaU1pMDVNalF5TFRka09USXlaVGRqTkdGak15SXNJbWx3WVdSa2NpSTZJakU0Tmk0eE1pNHhPRGN1TWpZaUxDSnVZVzFsSWpvaVZHVnpkRTFoY21OdmMweGhZbk1pTENKdmFXUWlPaUpsTTJNMU5XVmtZeTFqTXpRNExUUTBNVFl0T0dZd05TMHlOVFJtWmpNd05qVmpOV1VpTENKd2QyUmZkWEpzSWpvaWFIUjBjSE02THk5d2IzSjBZV3d1YldsamNtOXpiMlowYjI1c2FXNWxMbU52YlM5RGFHRnVaMlZRWVhOemQyOXlaQzVoYzNCNElpd2ljbWdpT2lJd0xrRldTVUU0T0ZSc0xXaHFlbWN3VXpoaU0xZFdXREJ2UzJOdFZGRXpTbHB1ZUUxa1QzQTNUbVZVVm5OV2FYVkhOa0ZRYnk0aUxDSnpZM0FpT2lKdFpHMWZaR1ZzWldkaGRHbHZiaUlzSW5OMVlpSTZJa1pTUTJ4RldURk9ObXR2ZEdWblMzcFplV0pFTjJkdFdGbGxhVTVIUkZrd05FSjJOV3R6ZDJGeGJVRWlMQ0owYVdRaU9pSm1ZV1UxWXpSbU15MW1NekU0TFRRME9ETXRZbU0yWmkwM05UazFOV1kwWVRCaE56SWlMQ0oxYm1seGRXVmZibUZ0WlNJNkluUmxjM1JBYldGeVkyOXpiR0ZpY3k1dmNtY2lMQ0oxY0c0aU9pSjBaWE4wUUcxaGNtTnZjMnhoWW5NdWIzSm5JaXdpZFhScElqb2lNVGg2WkVWSU5UZFRSWFZyYWpseGJqRm9aMlJCUVNJc0luWmxjaUk2SWpFdU1DSjkuVG1FUlRsZktBdWo5bTVvQUc2UTBRblV4VEFEaTNFamtlNHZ3VXo3UTdqUUFVZVZGZzl1U0pzUXNjU2hFTXVxUmQzN1R2VlpQanljdEVoRFgwLVpQcEVVYUlSempuRVEyTWxvc21SZURYZzhrYkhNZVliWi1jb0ZucDEyQkVpQnpJWFBGZnBpaU1GRnNZZ0hSSF9tSWxwYlBlRzJuQ2p0LTZSOHgzYVA5QS1tM0J3eV91dnV0WDFNVEVZRmFsekhGa04wNWkzbjZRcjhURnlJQ1ZUYW5OanlkMjBBZFRMbHJpTVk0RVBmZzRaLThVVTctZkcteElycWVPUmVWTnYwOUFHV192MDd6UkVaNmgxVk9tNl9nelRGcElVVURuZFdabnFLTHlySDlkdkF3WnFFSG1HUmlTNElNWnRFdDJNTkVZSnhDWHhlSi1VbWZJdV9tUVhKMW9R" + requestBytes, err := s.newSecurityTokenMsg(azureADTok, false, false) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2EnrollPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid RequestSecurityTokenResponseCollection message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("RequestSecurityTokenResponseCollection", resSoapMsg)) + require.True(t, s.isXMLTagPresent("DispositionMessage", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("TokenType", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("RequestID", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("BinarySecurityToken", resSoapMsg)) + + // Checking if an activity was created for the enrollment + s.lastActivityOfTypeMatches( + fleet.ActivityTypeMDMEnrolled{}.ActivityName(), + `{ + "mdm_platform": "microsoft", + "host_serial": "", + "installed_from_dep": false, + "host_display_name": "DESKTOP-0C89RC0" + }`, + 0) +} + +func (s *integrationMDMTestSuite) TestInvalidRequestSecurityTokenRequestWithMissingAdditionalContext() { + t := s.T() + + // create a new Host to get the UUID on the DB + windowsHost, err := s.ds.NewHost(context.Background(), &fleet.Host{ + ID: 1, + OsqueryHostID: ptr.String("Desktop-ABCQWE"), + NodeKey: ptr.String("Desktop-ABCQWE"), + UUID: uuid.New().String(), + Hostname: fmt.Sprintf("%sfoo.local.not.enrolled", s.T().Name()), + Platform: "windows", + }) + require.NoError(t, err) + + // Preparing the RequestSecurityToken Request message + encodedBinToken, err := GetEncodedBinarySecurityToken(fleet.WindowsMDMProgrammaticEnrollmentType, windowsHost.UUID) + require.NoError(t, err) + + requestBytes, err := s.newSecurityTokenMsg(encodedBinToken, true, true) + require.NoError(t, err) + + resp := s.DoRaw("POST", microsoft_mdm.MDE2EnrollPath, requestBytes, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + + require.Contains(t, resp.Header["Content-Type"], microsoft_mdm.SoapContentType) + + // Checking if SOAP response can be unmarshalled to an golang type + var xmlType interface{} + err = xml.Unmarshal(resBytes, &xmlType) + require.NoError(t, err) + + // Checking if SOAP response contains a valid SoapFault message + resSoapMsg := string(resBytes) + require.True(t, s.isXMLTagPresent("s:fault", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:value", resSoapMsg)) + require.True(t, s.isXMLTagContentPresent("s:text", resSoapMsg)) + require.True(t, s.checkIfXMLTagContains("s:text", "ContextItem item DeviceType is not present", resSoapMsg)) +} + +func (s *integrationMDMTestSuite) TestValidGetAuthRequest() { + t := s.T() + + // Target Endpoint url with query params + targetEndpointURL := microsoft_mdm.MDE2AuthPath + "?appru=ms-app%3A%2F%2Fwindows.immersivecontrolpanel&login_hint=demo%40mdmwindows.com" + resp := s.DoRaw("GET", targetEndpointURL, nil, http.StatusOK) + + resBytes, err := io.ReadAll(resp.Body) + require.NoError(t, err) + require.Contains(t, resp.Header["Content-Type"], "text/html; charset=UTF-8") + require.NotEmpty(t, resBytes) + + // Checking response content + resContent := string(resBytes) + require.Contains(t, resContent, "inputToken.name = 'wresult'") + require.Contains(t, resContent, "form.action = \"ms-app://windows.immersivecontrolpanel\"") + require.Contains(t, resContent, "performPost()") + + // Getting token content + encodedToken := s.getRawTokenValue(resContent) + require.NotEmpty(t, encodedToken) +} + +func (s *integrationMDMTestSuite) TestInvalidGetAuthRequest() { + t := s.T() + + // Target Endpoint url with no login_hit query param + targetEndpointURL := microsoft_mdm.MDE2AuthPath + "?appru=ms-app%3A%2F%2Fwindows.immersivecontrolpanel" + resp := s.DoRaw("GET", targetEndpointURL, nil, http.StatusInternalServerError) + + resBytes, err := io.ReadAll(resp.Body) + resContent := string(resBytes) + require.NoError(t, err) + require.NotEmpty(t, resBytes) + require.Contains(t, resContent, "forbidden") +} + // /////////////////////////////////////////////////////////////////////////// // Common helpers @@ -6020,6 +6181,24 @@ func (s *integrationMDMTestSuite) runWorker() { require.Empty(s.T(), pending) } +func (s *integrationMDMTestSuite) getRawTokenValue(content string) string { + // Create a regex object with the defined pattern + pattern := `inputToken.value\s*=\s*'([^']*)'` + regex := regexp.MustCompile(pattern) + + // Find the submatch using the regex pattern + submatches := regex.FindStringSubmatch(content) + + if len(submatches) >= 2 { + // Extract the content from the submatch + encodedToken := submatches[1] + + return encodedToken + } + + return "" +} + func (s *integrationMDMTestSuite) isXMLTagPresent(xmlTag string, payload string) bool { regex := fmt.Sprintf("<%s.*>", xmlTag) matched, err := regexp.MatchString(regex, payload) @@ -6051,11 +6230,17 @@ func (s *integrationMDMTestSuite) checkIfXMLTagContains(xmlTag string, xmlConten return true } -func (s *integrationMDMTestSuite) newGetPoliciesMsg(encodedBinToken string) ([]byte, error) { +func (s *integrationMDMTestSuite) newGetPoliciesMsg(deviceToken bool, encodedBinToken string) ([]byte, error) { if len(encodedBinToken) == 0 { return nil, errors.New("encodedBinToken is empty") } + // JWT token by default + tokType := microsoft_mdm.BinarySecurityAzureEnroll + if deviceToken { + tokType = microsoft_mdm.BinarySecurityDeviceEnroll + } + return []byte(` @@ -6066,7 +6251,7 @@ func (s *integrationMDMTestSuite) newGetPoliciesMsg(encodedBinToken string) ([]b https://mdmwindows.com/EnrollmentServer/Policy.svc - ` + encodedBinToken + ` + ` + encodedBinToken + ` @@ -6081,19 +6266,25 @@ func (s *integrationMDMTestSuite) newGetPoliciesMsg(encodedBinToken string) ([]b `), nil } -func (s *integrationMDMTestSuite) newSecurityTokenMsg(encodedBinToken string, missingContextItem bool) ([]byte, error) { +func (s *integrationMDMTestSuite) newSecurityTokenMsg(encodedBinToken string, deviceToken bool, missingContextItem bool) ([]byte, error) { if len(encodedBinToken) == 0 { return nil, errors.New("encodedBinToken is empty") } var reqSecTokenContextItemDeviceType []byte - if missingContextItem { + if !missingContextItem { reqSecTokenContextItemDeviceType = []byte( ` CIMClient_Windows `) } + // JWT token by default + tokType := microsoft_mdm.BinarySecurityAzureEnroll + if deviceToken { + tokType = microsoft_mdm.BinarySecurityDeviceEnroll + } + // Preparing the RequestSecurityToken Request message requestBytes := []byte( ` @@ -6105,7 +6296,7 @@ func (s *integrationMDMTestSuite) newSecurityTokenMsg(encodedBinToken string, mi https://mdmwindows.com/EnrollmentServer/Enrollment.svc - ` + encodedBinToken + ` + ` + encodedBinToken + ` diff --git a/server/service/microsoft_mdm.go b/server/service/microsoft_mdm.go index 54d8b9b998..7afe3abb68 100644 --- a/server/service/microsoft_mdm.go +++ b/server/service/microsoft_mdm.go @@ -1,6 +1,7 @@ package service import ( + "bytes" "context" "crypto/x509" "encoding/base64" @@ -8,8 +9,10 @@ import ( "encoding/xml" "errors" "fmt" + "html/template" "io" "net/http" + "net/url" "strconv" "time" @@ -23,22 +26,31 @@ import ( ) type SoapRequestContainer struct { - Data *fleet.SoapRequest - Err error + Data *fleet.SoapRequest + Params url.Values + Err error } // MDM SOAP request decoder -func (req *SoapRequestContainer) DecodeBody(ctx context.Context, r io.Reader) error { +func (req *SoapRequestContainer) DecodeBody(ctx context.Context, r io.Reader, u url.Values) error { // Reading the request bytes reqBytes, err := io.ReadAll(r) if err != nil { return ctxerr.Wrap(ctx, err, "reading soap mdm request") } - // Unmarshal the XML data from the request into the SoapRequest struct - err = xml.Unmarshal(reqBytes, &req.Data) - if err != nil { - return ctxerr.Wrap(ctx, err, "unmarshalling soap mdm request") + // Set the request parameters + req.Params = u + + // Handle empty body scenario + req.Data = &fleet.SoapRequest{} + + if len(reqBytes) != 0 { + // Unmarshal the XML data from the request into the SoapRequest struct + err = xml.Unmarshal(reqBytes, &req.Data) + if err != nil { + return ctxerr.Wrap(ctx, err, "unmarshalling soap mdm request") + } } return nil @@ -51,7 +63,7 @@ type SoapResponseContainer struct { func (r SoapResponseContainer) error() error { return r.Err } -// hijackRender writes the response header and the RAW XML output +// hijackRender writes the response header and the RAW HTML output func (r SoapResponseContainer) hijackRender(ctx context.Context, w http.ResponseWriter) { xmlRes, err := xml.MarshalIndent(r.Data, "", "\t") if err != nil { @@ -70,6 +82,23 @@ func (r SoapResponseContainer) hijackRender(ctx context.Context, w http.Response } } +type MDMAuthContainer struct { + Data *string + Err error +} + +func (r MDMAuthContainer) error() error { return r.Err } + +// hijackRender writes the response header and the RAW XML output +func (r MDMAuthContainer) hijackRender(ctx context.Context, w http.ResponseWriter) { + w.Header().Set("Content-Type", "text/html; charset=UTF-8") + w.Header().Set("Content-Length", strconv.Itoa(len(*r.Data))) + w.WriteHeader(http.StatusOK) + if n, err := w.Write([]byte(*r.Data)); err != nil { + logging.WithExtras(ctx, "err", err, "written", n) + } +} + // getUtcTime returns the current timestamp plus the specified number of minutes, // formatted as "2006-01-02T15:04:05.000Z". func getUtcTime(minutes int) string { @@ -82,7 +111,7 @@ func getUtcTime(minutes int) string { } // NewDiscoverResponse creates a new DiscoverResponse struct based on the auth policy, policy url, and enrollment url -func NewDiscoverResponse(authPolicy string, policyUrl string, enrollmentUrl string) (mdm_types.DiscoverResponse, error) { +func NewDiscoverResponse(authPolicy string, policyUrl string, enrollmentUrl string, authUrl *string) (mdm_types.DiscoverResponse, error) { if (len(authPolicy) == 0) || (len(policyUrl) == 0) || (len(enrollmentUrl) == 0) { return mdm_types.DiscoverResponse{}, errors.New("invalid parameters") } @@ -94,6 +123,7 @@ func NewDiscoverResponse(authPolicy string, policyUrl string, enrollmentUrl stri EnrollmentVersion: mdm.EnrollmentVersionV4, EnrollmentPolicyServiceUrl: policyUrl, EnrollmentServiceUrl: enrollmentUrl, + AuthServiceUrl: authUrl, }, }, nil } @@ -263,6 +293,14 @@ func NewSoapFault(errorType string, origMessage int, errorMessage error) mdm_typ } } +// getSTSAuthContent Retuns STS auth content +func getSTSAuthContent(data string) errorer { + return MDMAuthContainer{ + Data: &data, + Err: nil, + } +} + // getSoapResponseFault Returns a SoapResponse with a SoapFault on its body func getSoapResponseFault(relatesTo string, soapFault *mdm_types.SoapFault) errorer { if len(relatesTo) == 0 { @@ -408,11 +446,19 @@ func NewBinarySecurityTokenPayload(encodedToken string) (fleet.WindowsMDMAccessT return tokenPayload, nil } -// GetEncodedBinarySecurityToken returns the base64 form of a BinarySecurityTokenPayload -func GetEncodedBinarySecurityToken(typeID fleet.WindowsMDMEnrollmentType, hostUUID string) (string, error) { +// GetEncodedBinarySecurityToken returns the base64 form of a input payload +func GetEncodedBinarySecurityToken(typeID fleet.WindowsMDMEnrollmentType, payload string) (string, error) { var pld fleet.WindowsMDMAccessTokenPayload pld.Type = typeID - pld.Payload.HostUUID = hostUUID + + if typeID == fleet.WindowsMDMProgrammaticEnrollmentType { + pld.Payload.HostUUID = payload + } else if typeID == fleet.WindowsMDMAutomaticEnrollmentType { + pld.Payload.AuthToken = payload + } else { + return "", fmt.Errorf("invalid enrollment type: %v", typeID) + } + rawBytes, err := json.Marshal(pld) if err != nil { return "", err @@ -591,7 +637,7 @@ func mdmMicrosoftDiscoveryEndpoint(ctx context.Context, request interface{}, svc } // Getting the DiscoveryResponse message - discoveryResponseMsg, err := svc.GetMDMMicrosoftDiscoveryResponse(ctx) + discoveryResponseMsg, err := svc.GetMDMMicrosoftDiscoveryResponse(ctx, req.Body.Discover.Request.EmailAddress) if err != nil { soapFault := svc.GetAuthorizedSoapFault(ctx, mdm.SoapErrorMessageFormat, mdm_types.MDEDiscovery, err) return getSoapResponseFault(req.GetMessageID(), soapFault), nil @@ -610,6 +656,31 @@ func mdmMicrosoftDiscoveryEndpoint(ctx context.Context, request interface{}, svc }, nil } +// mdmMicrosoftAuthEndpoint handles the Security Token Service (STS) implementation +func mdmMicrosoftAuthEndpoint(ctx context.Context, request interface{}, svc fleet.Service) (errorer, error) { + params := request.(*SoapRequestContainer).Params + + // Sanity check on the expected query params + if !params.Has(mdm.STSAuthAppRu) || !params.Has(mdm.STSLoginHint) { + return getSTSAuthContent(""), errors.New("expected STS params are not present") + } + + appru := params.Get(mdm.STSAuthAppRu) + loginHint := params.Get(mdm.STSLoginHint) + + if (len(appru) == 0) || (len(loginHint) == 0) { + return getSTSAuthContent(""), errors.New("expected STS params are empty") + } + + // Getting the STS endpoint HTML content + stsAuthContent, err := svc.GetMDMMicrosoftSTSAuthResponse(ctx, appru, loginHint) + if err != nil { + return getSTSAuthContent(""), errors.New("error generating STS content") + } + + return getSTSAuthContent(stsAuthContent), nil +} + // mdmMicrosoftPolicyEndpoint handles the GetPolicies message and returns a valid GetPoliciesResponse message // GetPoliciesResponse message contains the certificate policies required for the next enrollment step. For more information about these messages, see [MS-XCEP] sections 3.1.4.1.1.1 and 3.1.4.1.1.2. func mdmMicrosoftPolicyEndpoint(ctx context.Context, request interface{}, svc fleet.Service) (errorer, error) { @@ -622,14 +693,14 @@ func mdmMicrosoftPolicyEndpoint(ctx context.Context, request interface{}, svc fl } // Binary security token should be extracted to ensure this is a valid call - binSecTokenData, err := req.GetBinarySecurityToken() + hdrSecToken, err := req.GetHeaderBinarySecurityToken() if err != nil { soapFault := svc.GetAuthorizedSoapFault(ctx, mdm.SoapErrorMessageFormat, mdm_types.MDEPolicy, err) return getSoapResponseFault(req.GetMessageID(), soapFault), nil } // Getting the GetPoliciesResponse message - policyResponseMsg, err := svc.GetMDMWindowsPolicyResponse(ctx, binSecTokenData) + policyResponseMsg, err := svc.GetMDMWindowsPolicyResponse(ctx, hdrSecToken) if err != nil { soapFault := svc.GetAuthorizedSoapFault(ctx, mdm.SoapErrorMessageFormat, mdm_types.MDEPolicy, err) return getSoapResponseFault(req.GetMessageID(), soapFault), nil @@ -667,14 +738,14 @@ func mdmMicrosoftEnrollEndpoint(ctx context.Context, request interface{}, svc fl } // Binary security token should be extracted to ensure this is a valid call - binSecTokenData, err := req.GetBinarySecurityToken() + hdrBinarySecToken, err := req.GetHeaderBinarySecurityToken() if err != nil { soapFault := svc.GetAuthorizedSoapFault(ctx, mdm.SoapErrorMessageFormat, mdm_types.MDEEnrollment, err) return getSoapResponseFault(req.GetMessageID(), soapFault), nil } // Getting the RequestSecurityTokenResponseCollection message - enrollResponseMsg, err := svc.GetMDMWindowsEnrollResponse(ctx, reqSecurityTokenMsg, binSecTokenData) + enrollResponseMsg, err := svc.GetMDMWindowsEnrollResponse(ctx, reqSecurityTokenMsg, hdrBinarySecToken) if err != nil { soapFault := svc.GetAuthorizedSoapFault(ctx, mdm.SoapErrorMessageFormat, mdm_types.MDEEnrollment, err) return getSoapResponseFault(req.GetMessageID(), soapFault), nil @@ -693,45 +764,79 @@ func mdmMicrosoftEnrollEndpoint(ctx context.Context, request interface{}, svc fl }, nil } -// validateBinarySecurityToken checks if the provided token is valid -func (svc *Service) validateBinarySecurityToken(ctx context.Context, encodedBinarySecToken string) error { - if len(encodedBinarySecToken) == 0 { - return errors.New("binarySecurityTokenValidation: encoded token is invalid") +// authBinarySecurityToken checks if the provided token is valid +func (svc *Service) authBinarySecurityToken(ctx context.Context, authToken *fleet.HeaderBinarySecurityToken) (string, error) { + if authToken == nil { + return "", errors.New("authToken is empty") } - // Getting the Binary Security Token Payload - binSecToken, err := NewBinarySecurityTokenPayload(encodedBinarySecToken) + err := authToken.IsValidToken() if err != nil { - return fmt.Errorf("binarySecurityTokenValidation: token creation error %v", err) + return "", errors.New("authToken is not valid") } - // Validating the Binary Security Token Payload - err = binSecToken.IsValidToken() - if err != nil { - return fmt.Errorf("binarySecurityTokenValidation: invalid token data %v", err) - } + // Tokens that were generated by enrollment client + if authToken.IsDeviceToken() { - // Validating the Binary Security Token Type used on Programmatic Enrollments - if binSecToken.Type == mdm_types.WindowsMDMProgrammaticEnrollmentType { - host, err := svc.ds.HostByIdentifier(ctx, binSecToken.Payload.HostUUID) + // Getting the Binary Security Token Payload + binSecToken, err := NewBinarySecurityTokenPayload(authToken.Content) if err != nil { - return fmt.Errorf("binarySecurityTokenValidation: host data cannot be found %v", err) + return "", fmt.Errorf("token creation error %v", err) } - // This ensures that only hosts that are eligible for Windows enrollment can be enrolled - if !host.IsEligibleForWindowsMDMEnrollment() { - return errors.New("binarySecurityTokenValidation: host is not elegible for Windows MDM enrollment") + // Validating the Binary Security Token Payload + err = binSecToken.IsValidToken() + if err != nil { + return "", fmt.Errorf("invalid token data %v", err) + } + + // Validating the Binary Security Token Type used on Programmatic Enrollments + if binSecToken.Type == mdm_types.WindowsMDMProgrammaticEnrollmentType { + host, err := svc.ds.HostByIdentifier(ctx, binSecToken.Payload.HostUUID) + if err != nil { + return "", fmt.Errorf("host data cannot be found %v", err) + } + + // This ensures that only hosts that are eligible for Windows enrollment can be enrolled + if !host.IsEligibleForWindowsMDMEnrollment() { + return "", errors.New("host is not elegible for Windows MDM enrollment") + } + + // No errors, token is authorized + return binSecToken.Payload.HostUUID, nil + } + + // Validating the Binary Security Token Type used on Automatic Enrollments (returned by STS Auth Endpoint) + if binSecToken.Type == mdm_types.WindowsMDMAutomaticEnrollmentType { + + upnToken, err := svc.wstepCertManager.GetSTSAuthTokenUPNClaim(binSecToken.Payload.AuthToken) + if err != nil { + return "", ctxerr.Wrap(ctx, err, "issue retrieving UPN from Auth token") + } + + // No errors, token is authorized + return upnToken, nil + } + } + + // Validating the Binary Security Token Type used on Automatic Enrollments + if authToken.IsAzureJWTToken() { + + // Validate the JWT Auth token by retreving its claims + tokenData, err := mdm.GetAzureAuthTokenClaims(authToken.Content) + if err != nil { + return "", fmt.Errorf("binary security token claim failed: %v", err) } // No errors, token is authorized - return nil + return tokenData.UPN, nil } - return errors.New("binarySecurityTokenValidation: token is not authorized") + return "", errors.New("token is not authorized") } // GetMDMMicrosoftDiscoveryResponse returns a valid DiscoveryResponse message -func (svc *Service) GetMDMMicrosoftDiscoveryResponse(ctx context.Context) (*fleet.DiscoverResponse, error) { +func (svc *Service) GetMDMMicrosoftDiscoveryResponse(ctx context.Context, upnEmail string) (*fleet.DiscoverResponse, error) { // skipauth: This endpoint does not use authentication svc.authz.SkipAuthorization(ctx) @@ -752,7 +857,18 @@ func (svc *Service) GetMDMMicrosoftDiscoveryResponse(ctx context.Context) (*flee return nil, ctxerr.Wrap(ctx, err, "resolve enroll endpoint") } - discoveryMsg, err := NewDiscoverResponse(mdm.AuthOnPremise, urlPolicyEndpoint, urlEnrollEndpoint) + // Only adding STS Auth endpoint if the UPN email is provided + var urlSTSAuthEndpoint *string + if len(upnEmail) > 0 { + workUrlSTSAuthEndpoint, err := mdm.ResolveWindowsMDMAuth(appCfg.ServerSettings.ServerURL) + if err != nil { + return nil, ctxerr.Wrap(ctx, err, "resolve enroll endpoint") + } + + urlSTSAuthEndpoint = &workUrlSTSAuthEndpoint + } + + discoveryMsg, err := NewDiscoverResponse(mdm.AuthOnPremise, urlPolicyEndpoint, urlEnrollEndpoint, urlSTSAuthEndpoint) if err != nil { return nil, ctxerr.Wrap(ctx, err, "creation of DiscoverResponse message") } @@ -760,14 +876,68 @@ func (svc *Service) GetMDMMicrosoftDiscoveryResponse(ctx context.Context) (*flee return &discoveryMsg, nil } +// GetMDMMicrosoftSTSAuthResponse returns a valid Security Token Service (STS) page content +func (svc *Service) GetMDMMicrosoftSTSAuthResponse(ctx context.Context, appru string, loginHint string) (string, error) { + // skipauth: This endpoint does not use authentication + svc.authz.SkipAuthorization(ctx) + + // Dummy data will be returned as part of the token as user-driven enrollment is not supported yet + // In the future, the following calls would have to be made to support user-driven enrollment + // encodedBST will carry the token to return + // authToken, err := svc.wstepCertManager.NewSTSAuthToken(loginHint) + // encodedBST, err := GetEncodedBinarySecurityToken(fleet.WindowsMDMAutomaticEnrollmentType, authToken) + encodedBST := "user_driven_enrollment_not_implemented" + + // STS Auth Endpoint returns HTML content that gets render in a webview container + // The webview container expect a POST request to the appru URL with the wresult parameter set to the auth token + // The security token in wresult is later passed back in + // This string is opaque to the enrollment client; the client does not interpret the string. + // The returned HTML content contains a JS script that will perform a POST request to the appru URL automatically + // This will set the wresult parameter to the value of auth token + tmpl, err := template.New("").Parse(` + + `) + if err != nil { + return "", ctxerr.Wrap(ctx, err, "STS content template") + } + + var htmlBuf bytes.Buffer + err = tmpl.Execute(&htmlBuf, map[string][]byte{"ActionURL": []byte(appru), "Token": []byte(encodedBST)}) + if err != nil { + return "", ctxerr.Wrap(ctx, err, "creation of STS content") + } + + return htmlBuf.String(), nil +} + // GetMDMWindowsPolicyResponse returns a valid GetPoliciesResponse message -func (svc *Service) GetMDMWindowsPolicyResponse(ctx context.Context, authToken string) (*fleet.GetPoliciesResponse, error) { - if len(authToken) == 0 { - return nil, fleet.NewInvalidArgumentError("policy response", "authToken is empty") +func (svc *Service) GetMDMWindowsPolicyResponse(ctx context.Context, authToken *fleet.HeaderBinarySecurityToken) (*fleet.GetPoliciesResponse, error) { + if authToken == nil { + return nil, fleet.NewInvalidArgumentError("policy response", "authToken is invalid") } // Validate the binary security token - err := svc.validateBinarySecurityToken(ctx, authToken) + _, err := svc.authBinarySecurityToken(ctx, authToken) if err != nil { return nil, ctxerr.Wrap(ctx, err, "validate binary security token") } @@ -787,13 +957,13 @@ func (svc *Service) GetMDMWindowsPolicyResponse(ctx context.Context, authToken s // GetMDMWindowsEnrollResponse returns a valid RequestSecurityTokenResponseCollection message // secTokenMsg is the RequestSecurityToken message // authToken is the base64 encoded binary security token -func (svc *Service) GetMDMWindowsEnrollResponse(ctx context.Context, secTokenMsg *fleet.RequestSecurityToken, authToken string) (*fleet.RequestSecurityTokenResponseCollection, error) { - if len(authToken) == 0 { - return nil, fleet.NewInvalidArgumentError("enroll response", "authToken is empty") +func (svc *Service) GetMDMWindowsEnrollResponse(ctx context.Context, secTokenMsg *fleet.RequestSecurityToken, authToken *fleet.HeaderBinarySecurityToken) (*fleet.RequestSecurityTokenResponseCollection, error) { + if authToken == nil { + return nil, fleet.NewInvalidArgumentError("enroll response", "authToken is not present") } - // Validate the binary security token - err := svc.validateBinarySecurityToken(ctx, authToken) + // Auth the binary security token + userID, err := svc.authBinarySecurityToken(ctx, authToken) if err != nil { return nil, ctxerr.Wrap(ctx, err, "validate binary security token") } @@ -828,7 +998,7 @@ func (svc *Service) GetMDMWindowsEnrollResponse(ctx context.Context, secTokenMsg // // This method also creates the relevant enrollment activity as it has // access to the device information. - err = svc.storeWindowsMDMEnrolledDevice(ctx, secTokenMsg) + err = svc.storeWindowsMDMEnrolledDevice(ctx, userID, secTokenMsg) if err != nil { return nil, ctxerr.Wrap(ctx, err, "enrolled device information cannot be stored") } @@ -944,7 +1114,7 @@ func (svc *Service) getDeviceProvisioningInformation(ctx context.Context, secTok } // storeWindowsMDMEnrolledDevice stores the device information to the list of MDM enrolled devices -func (svc *Service) storeWindowsMDMEnrolledDevice(ctx context.Context, secTokenMsg *fleet.RequestSecurityToken) error { +func (svc *Service) storeWindowsMDMEnrolledDevice(ctx context.Context, userID string, secTokenMsg *fleet.RequestSecurityToken) error { const ( error_tag = "windows MDM enrolled storage: " ) @@ -999,7 +1169,7 @@ func (svc *Service) storeWindowsMDMEnrolledDevice(ctx context.Context, secTokenM MDMDeviceType: reqDeviceType, MDMDeviceName: reqDeviceName, MDMEnrollType: reqEnrollType, - MDMEnrollUserID: "", // No user information is available at this point + MDMEnrollUserID: userID, // This could be Host UUID or UPN email MDMEnrollProtoVersion: reqEnrollVersion, MDMEnrollClientVersion: reqAppVersion, MDMNotInOOBE: false, diff --git a/server/service/teams.go b/server/service/teams.go index ee9dd0da04..0356616295 100644 --- a/server/service/teams.go +++ b/server/service/teams.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "net/http" + "net/url" "github.com/fleetdm/fleet/v4/server/contexts/ctxerr" "github.com/fleetdm/fleet/v4/server/fleet" @@ -180,7 +181,7 @@ type applyTeamSpecsRequest struct { Specs []*fleet.TeamSpec `json:"specs"` } -func (req *applyTeamSpecsRequest) DecodeBody(ctx context.Context, r io.Reader) error { +func (req *applyTeamSpecsRequest) DecodeBody(ctx context.Context, r io.Reader, u url.Values) error { if err := fleet.JSONStrictDecode(r, req); err != nil { err = fleet.NewUserMessageError(err, http.StatusBadRequest) if !req.Force || !fleet.IsJSONUnknownFieldError(err) {