* clear

* feat: Add everything

* fix: Commenht
This commit is contained in:
Prad Nukala
2025-10-03 14:45:52 -04:00
committed by GitHub
parent 43b4a11c06
commit 13e6c3e84d
1935 changed files with 655061 additions and 40058 deletions
+265
View File
@@ -0,0 +1,265 @@
// Package handlers provides HTTP handlers for the highway server
package handlers
import (
"crypto/rand"
"encoding/base64"
"fmt"
"net/http"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/labstack/echo/v4"
)
// CustomClaims defines custom JWT claims for vault operations
type CustomClaims struct {
UserID string `json:"user_id"`
Permissions []string `json:"permissions"`
jwt.RegisteredClaims
}
// LoginRequest represents the login request payload
type LoginRequest struct {
Username string `json:"username"`
Password string `json:"password"`
}
// LoginResponse represents the login response payload
type LoginResponse struct {
Token string `json:"token"`
}
// LoginHandler generates JWT tokens for authentication
// This now supports both traditional login and OIDC-based authentication
func LoginHandler(jwtSecret []byte) echo.HandlerFunc {
return func(c echo.Context) error {
// Check if this is an OIDC callback
code := c.QueryParam("code")
if code != "" {
return handleOIDCCallback(c, code, jwtSecret)
}
// Check if user has WebAuthn credentials
authHeader := c.Request().Header.Get("X-WebAuthn-Assertion")
if authHeader != "" {
return handleWebAuthnLogin(c, authHeader, jwtSecret)
}
// Traditional login flow
var req LoginRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
// Simple authentication - in production, validate against a proper user database
if req.Username == "" || req.Password == "" {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Username and password are required"},
)
}
// For demo purposes, accept any non-empty credentials
// In production, verify credentials against database/directory
if req.Username == "vault-user" && req.Password == "vault-pass" {
// Create custom claims
claims := &CustomClaims{
UserID: req.Username,
Permissions: []string{
"vault:generate",
"vault:sign",
"vault:verify",
"vault:export",
"vault:import",
"vault:refresh",
},
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour * 24)), // 24 hours
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
Issuer: "highway-vault",
Subject: req.Username,
},
}
// Create token with claims
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
// Sign token with secret
tokenString, err := token.SignedString(jwtSecret)
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to generate token"},
)
}
return c.JSON(http.StatusOK, LoginResponse{Token: tokenString})
}
return c.JSON(http.StatusUnauthorized, map[string]string{"error": "Invalid credentials"})
}
}
// handleOIDCCallback processes OIDC authorization code callback
func handleOIDCCallback(c echo.Context, code string, jwtSecret []byte) error {
// Exchange code for tokens using OIDC token endpoint
tokenReq := &OIDCTokenRequest{
GrantType: "authorization_code",
Code: code,
RedirectURI: c.QueryParam("redirect_uri"),
ClientID: c.QueryParam("client_id"),
}
// Call internal OIDC token handler
// In production, this would make an HTTP call to the OIDC provider
c.Set("oidc_token_request", tokenReq)
if err := handleAuthorizationCodeGrant(c, tokenReq); err != nil {
return err
}
// Get the OIDC session from context
session, ok := c.Get("oidc_session").(*OIDCSession)
if !ok {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "Failed to establish OIDC session",
})
}
// Create JWT token from OIDC session
claims := &CustomClaims{
UserID: session.UserDID,
Permissions: []string{
"vault:generate",
"vault:sign",
"vault:verify",
"vault:export",
"vault:import",
"vault:refresh",
},
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(session.ExpiresAt),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
Issuer: "highway-vault-oidc",
Subject: session.UserDID,
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
tokenString, err := token.SignedString(jwtSecret)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "Failed to generate token",
})
}
return c.JSON(http.StatusOK, LoginResponse{Token: tokenString})
}
// handleWebAuthnLogin processes WebAuthn-based login
func handleWebAuthnLogin(c echo.Context, assertion string, jwtSecret []byte) error {
// Verify WebAuthn assertion
// This would integrate with the WebAuthn handlers
// For now, extract username from assertion (simplified)
username := c.Request().Header.Get("X-WebAuthn-Username")
if username == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "WebAuthn username required",
})
}
// Verify the assertion matches a stored credential
webAuthnStore.mu.RLock()
credentials, exists := webAuthnStore.credentials[username]
webAuthnStore.mu.RUnlock()
if !exists || len(credentials) == 0 {
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "No WebAuthn credentials found for user",
})
}
// Create JWT token for WebAuthn authenticated user
claims := &CustomClaims{
UserID: username,
Permissions: []string{
"vault:generate",
"vault:sign",
"vault:verify",
"vault:export",
"vault:import",
"vault:refresh",
},
RegisteredClaims: jwt.RegisteredClaims{
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour * 24)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
Issuer: "highway-vault-webauthn",
Subject: username,
},
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
tokenString, err := token.SignedString(jwtSecret)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "Failed to generate token",
})
}
return c.JSON(http.StatusOK, LoginResponse{Token: tokenString})
}
// OIDCLoginHandler initiates OIDC login flow
func OIDCLoginHandler() echo.HandlerFunc {
return func(c echo.Context) error {
// Build authorization URL
authURL := buildOIDCAuthorizationURL(
c.QueryParam("client_id"),
c.QueryParam("redirect_uri"),
c.QueryParam("scope"),
c.QueryParam("state"),
)
// Redirect to OIDC authorization endpoint
return c.Redirect(http.StatusFound, authURL)
}
}
// buildOIDCAuthorizationURL constructs the OIDC authorization URL
func buildOIDCAuthorizationURL(clientID, redirectURI, scope, state string) string {
// Default values if not provided
if clientID == "" {
clientID = "highway-vault-client"
}
if redirectURI == "" {
redirectURI = "http://localhost:8080/auth/callback"
}
if scope == "" {
scope = "openid profile did vault"
}
if state == "" {
state = generateState()
}
// Build authorization URL
return fmt.Sprintf(
"https://localhost:8080/oidc/authorize?response_type=code&client_id=%s&redirect_uri=%s&scope=%s&state=%s",
clientID,
redirectURI,
scope,
state,
)
}
// generateState generates a random state parameter for OIDC
func generateState() string {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return "default-state"
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
+60
View File
@@ -0,0 +1,60 @@
package handlers
import (
"encoding/json"
"testing"
"github.com/sonr-io/sonr/crypto/mpc"
)
// getQueueFromPriority returns the queue name based on priority
func getQueueFromPriority(priority string) string {
return GetQueueFromPriority(priority)
}
// BenchmarkJSONMarshaling measures JSON encoding/decoding performance
func BenchmarkJSONMarshaling(b *testing.B) {
payload := map[string]any{
"message": []byte("benchmark test message for JSON marshaling performance"),
"enclave": &mpc.EnclaveData{},
"priority": "critical",
}
b.ResetTimer()
b.ReportAllocs()
b.RunParallel(func(pb *testing.PB) {
for pb.Next() {
// Marshal
data, err := json.Marshal(payload)
if err != nil {
b.Error(err)
}
// Unmarshal
var decoded map[string]any
err = json.Unmarshal(data, &decoded)
if err != nil {
b.Error(err)
}
}
})
}
// BenchmarkQueuePrioritySelection measures queue selection performance
func BenchmarkQueuePrioritySelection(b *testing.B) {
priorities := []string{"critical", "high", "default", "low", "", "unknown"}
b.ResetTimer()
b.ReportAllocs()
b.RunParallel(func(pb *testing.PB) {
i := 0
for pb.Next() {
priority := priorities[i%len(priorities)]
i++
queue := getQueueFromPriority(priority)
_ = queue // Avoid compiler optimization
}
})
}
+333
View File
@@ -0,0 +1,333 @@
package handlers
import (
"fmt"
"net/http"
"time"
"github.com/labstack/echo/v4"
)
// BroadcastService handles blockchain broadcasting operations
type BroadcastService struct {
// TODO: Add cosmos SDK client for actual broadcasting
}
var broadcastService = &BroadcastService{}
// HandleBroadcast handles generic message broadcasting to blockchain
func HandleBroadcast(c echo.Context) error {
var req BroadcastRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Invalid broadcast request",
})
}
// Validate request
if req.Message == nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Message is required",
})
}
// Process based on gasless flag
if req.Gasless {
return handleGaslessBroadcast(c, &req)
}
return handleStandardBroadcast(c, &req)
}
// handleGaslessBroadcast handles gasless transaction broadcasting
func handleGaslessBroadcast(c echo.Context, req *BroadcastRequest) error {
// Validate gasless eligibility
if !isGaslessEligible(req.Message) {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Message not eligible for gasless broadcast",
})
}
// TODO: Implement actual blockchain broadcast
// For now, simulate successful broadcast
response := &BroadcastResponse{
TxHash: generateTxHash(),
Height: 12345,
Code: 0,
RawLog: "Transaction broadcast successfully",
Success: true,
}
return c.JSON(http.StatusOK, response)
}
// handleStandardBroadcast handles regular transaction broadcasting
func handleStandardBroadcast(c echo.Context, req *BroadcastRequest) error {
// Get sender address
fromAddress := req.FromAddress
if fromAddress == "" {
// Try to get from context
userDID := c.Get("user_did")
if userDID != nil {
fromAddress = deriveAddressFromDID(userDID.(string))
}
}
if fromAddress == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "From address is required for standard broadcast",
})
}
// TODO: Implement actual blockchain broadcast
// For now, simulate successful broadcast
response := &BroadcastResponse{
TxHash: generateTxHash(),
Height: 12346,
Code: 0,
RawLog: "Transaction broadcast successfully",
Success: true,
}
return c.JSON(http.StatusOK, response)
}
// BroadcastWebAuthnRegistration broadcasts WebAuthn registration to blockchain
func BroadcastWebAuthnRegistration(
credential *WebAuthnCredential,
gasless bool,
) (*BroadcastResponse, error) {
// Create MsgRegisterWebAuthnCredential
msg := map[string]any{
"@type": "/sonr.did.v1.MsgRegisterWebAuthnCredential",
"username": credential.Username,
"credential_id": credential.CredentialID,
"public_key": credential.PublicKey,
"attestation_object": credential.AttestationObject,
"client_data_json": credential.ClientDataJSON,
"origin": credential.Origin,
"algorithm": credential.Algorithm,
}
// Create broadcast request (for future use with actual broadcast)
_ = &BroadcastRequest{
Message: msg,
Gasless: gasless,
AutoSign: true,
}
// Simulate broadcast (TODO: Implement actual broadcast)
return &BroadcastResponse{
TxHash: generateTxHash(),
Height: 12347,
Code: 0,
RawLog: "WebAuthn credential registered successfully",
Success: true,
}, nil
}
// BroadcastVaultCreation broadcasts vault creation to blockchain
func BroadcastVaultCreation(
userDID string,
vaultConfig map[string]any,
) (*BroadcastResponse, error) {
// Create MsgCreateVault
msg := map[string]any{
"@type": "/sonr.vault.v1.MsgCreateVault",
"creator": userDID,
"config": vaultConfig,
}
// Create broadcast request (for future use with actual broadcast)
_ = &BroadcastRequest{
Message: msg,
Gasless: true, // Vault creation is gasless for new users
AutoSign: true,
FromAddress: deriveAddressFromDID(userDID),
}
// Simulate broadcast (TODO: Implement actual broadcast)
return &BroadcastResponse{
TxHash: generateTxHash(),
Height: 12348,
Code: 0,
RawLog: "Vault created successfully",
Success: true,
}, nil
}
// BroadcastDIDDocument broadcasts DID document to blockchain
func BroadcastDIDDocument(didDoc map[string]any, gasless bool) (*BroadcastResponse, error) {
// Create MsgCreateDIDDocument or MsgUpdateDIDDocument
msg := map[string]any{
"@type": "/sonr.did.v1.MsgCreateDIDDocument",
"did_document": didDoc,
}
// Create broadcast request (for future use with actual broadcast)
_ = &BroadcastRequest{
Message: msg,
Gasless: gasless,
AutoSign: true,
}
// Simulate broadcast (TODO: Implement actual broadcast)
return &BroadcastResponse{
TxHash: generateTxHash(),
Height: 12349,
Code: 0,
RawLog: "DID document created successfully",
Success: true,
}, nil
}
// HandleTransactionStatus checks transaction status
func HandleTransactionStatus(c echo.Context) error {
txHash := c.Param("hash")
if txHash == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Transaction hash is required",
})
}
// TODO: Query actual blockchain for transaction status
// For now, return simulated status
status := map[string]any{
"tx_hash": txHash,
"height": 12350,
"status": "confirmed",
"code": 0,
"gas_used": 50000,
"gas_wanted": 100000,
"timestamp": time.Now().Unix(),
}
return c.JSON(http.StatusOK, status)
}
// HandleEstimateGas estimates gas for a transaction
func HandleEstimateGas(c echo.Context) error {
var req BroadcastRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Invalid request",
})
}
// Check if gasless eligible
if isGaslessEligible(req.Message) {
return c.JSON(http.StatusOK, map[string]any{
"gas_estimate": 0,
"gasless": true,
"fee": "0usnr",
})
}
// TODO: Implement actual gas estimation
// For now, return default estimate
return c.JSON(http.StatusOK, map[string]any{
"gas_estimate": 100000,
"gasless": false,
"fee": "100usnr",
})
}
// Helper functions
// isGaslessEligible checks if a message is eligible for gasless broadcasting
func isGaslessEligible(message any) bool {
// Check message type
msgMap, ok := message.(map[string]any)
if !ok {
return false
}
msgType, ok := msgMap["@type"].(string)
if !ok {
return false
}
// WebAuthn registration and vault creation are gasless
gaslessTypes := []string{
"/sonr.did.v1.MsgRegisterWebAuthnCredential",
"/sonr.vault.v1.MsgCreateVault",
"/sonr.did.v1.MsgCreateDIDDocument",
}
for _, t := range gaslessTypes {
if msgType == t {
return true
}
}
return false
}
// deriveAddressFromDID derives a blockchain address from a DID
func deriveAddressFromDID(did string) string {
// TODO: Implement actual address derivation
// For now, return a placeholder address
return "sonr1placeholder" + did[len(did)-10:]
}
// generateTxHash generates a mock transaction hash
func generateTxHash() string {
// TODO: Replace with actual tx hash from broadcast
return fmt.Sprintf("%X", time.Now().UnixNano())
}
// CreateWebAuthnBroadcastHandler creates a handler that broadcasts WebAuthn credentials
func CreateWebAuthnBroadcastHandler() echo.HandlerFunc {
return func(c echo.Context) error {
// Get WebAuthn credential from context
credential, ok := c.Get("webauthn_credential").(*WebAuthnCredential)
if !ok {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "No WebAuthn credential found",
})
}
// Broadcast to blockchain
response, err := BroadcastWebAuthnRegistration(credential, true)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": err.Error(),
})
}
return c.JSON(http.StatusOK, response)
}
}
// CreateVaultBroadcastHandler creates a handler that broadcasts vault creation
func CreateVaultBroadcastHandler() echo.HandlerFunc {
return func(c echo.Context) error {
// Get user DID from context
userDID, ok := c.Get("user_did").(string)
if !ok {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "User DID not found",
})
}
// Get vault config from request
var vaultConfig map[string]any
if err := c.Bind(&vaultConfig); err != nil {
// Use default config
vaultConfig = map[string]any{
"type": "standard",
"encryption": "AES256",
}
}
// Broadcast vault creation
response, err := BroadcastVaultCreation(userDID, vaultConfig)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": err.Error(),
})
}
return c.JSON(http.StatusOK, response)
}
}
+33
View File
@@ -0,0 +1,33 @@
package handlers
import "errors"
// Common errors for OAuth2 and UCAN handling
var (
// Client errors
ErrClientNotFound = errors.New("client not found")
ErrInvalidClientCredentials = errors.New("invalid client credentials")
// Token errors
ErrTokenNotFound = errors.New("token not found")
ErrTokenExpired = errors.New("token has expired")
ErrInvalidToken = errors.New("invalid token")
ErrTokenRevoked = errors.New("token has been revoked")
// Authorization errors
ErrUnauthorized = errors.New("unauthorized")
ErrInsufficientScope = errors.New("insufficient scope")
ErrInvalidScope = errors.New("invalid scope")
ErrScopeNotAllowed = errors.New("scope not allowed for client")
// UCAN errors
ErrInvalidAttenuation = errors.New("invalid attenuation")
ErrInvalidDelegation = errors.New("invalid delegation")
ErrBrokenChain = errors.New("broken delegation chain")
ErrPrivilegeEscalation = errors.New("privilege escalation attempted")
// OIDC errors
ErrInvalidRedirectURI = errors.New("invalid redirect URI")
ErrInvalidResponseType = errors.New("invalid response type")
ErrInvalidGrantType = errors.New("invalid grant type")
)
+27
View File
@@ -0,0 +1,27 @@
package handlers
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestGetQueueFromPriority(t *testing.T) {
tests := []struct {
priority string
expected string
}{
{"critical", "critical"},
{"high", "critical"},
{"low", "low"},
{"", "default"},
{"unknown", "default"},
}
for _, tt := range tests {
t.Run(tt.priority, func(t *testing.T) {
result := GetQueueFromPriority(tt.priority)
assert.Equal(t, tt.expected, result)
})
}
}
+201
View File
@@ -0,0 +1,201 @@
package handlers
import (
"context"
"net/http"
"sync"
"time"
"github.com/hibiken/asynq"
"github.com/labstack/echo/v4"
"github.com/sonr-io/sonr/types/ipfs"
)
// HealthStatus represents the health status of the service
type HealthStatus struct {
Status string `json:"status"`
Timestamp string `json:"timestamp"`
Uptime string `json:"uptime"`
Dependencies map[string]string `json:"dependencies"`
}
// HealthChecker manages health and readiness checks
type HealthChecker struct {
startTime time.Time
redisClient *asynq.Client
ipfsClient ipfs.IPFSClient
ready bool
readyMu sync.RWMutex
redisHealthy bool
ipfsHealthy bool
}
// NewHealthChecker creates a new health checker
func NewHealthChecker(redisClient *asynq.Client, ipfsClient ipfs.IPFSClient) *HealthChecker {
hc := &HealthChecker{
startTime: time.Now(),
redisClient: redisClient,
ipfsClient: ipfsClient,
ready: false,
}
// Start background health checks
go hc.startHealthChecks()
return hc
}
// startHealthChecks runs periodic health checks
func (hc *HealthChecker) startHealthChecks() {
// Initial startup delay
time.Sleep(3 * time.Second)
ticker := time.NewTicker(10 * time.Second)
defer ticker.Stop()
// Run initial check
hc.checkDependencies()
for range ticker.C {
hc.checkDependencies()
}
}
// checkDependencies checks all service dependencies
func (hc *HealthChecker) checkDependencies() {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
// Check Redis connectivity
hc.redisHealthy = hc.checkRedis(ctx)
// Check IPFS connectivity (optional)
hc.ipfsHealthy = hc.checkIPFS(ctx)
// Update readiness based on critical dependencies
hc.readyMu.Lock()
hc.ready = hc.redisHealthy // Redis is required
hc.readyMu.Unlock()
}
// checkRedis verifies Redis connectivity
func (hc *HealthChecker) checkRedis(ctx context.Context) bool {
if hc.redisClient == nil {
return false
}
// Try to ping Redis by enqueuing a test task
testTask := asynq.NewTask("health:check", nil)
_, err := hc.redisClient.EnqueueContext(ctx, testTask,
asynq.Queue("health"),
asynq.MaxRetry(0),
asynq.Retention(1*time.Second)) // Auto-delete after 1 second
if err != nil {
return false
}
return true
}
// checkIPFS verifies IPFS connectivity
func (hc *HealthChecker) checkIPFS(ctx context.Context) bool {
if hc.ipfsClient == nil {
return true // IPFS is optional
}
// Check IPFS connectivity by getting version
ch := make(chan bool, 1)
go func() {
// Try a simple operation to check connectivity
if hc.ipfsClient != nil {
// IPFS client exists, assume healthy for now
// Real check would depend on actual IPFS client implementation
ch <- true
} else {
ch <- false
}
}()
select {
case result := <-ch:
return result
case <-ctx.Done():
return false
}
}
// IsReady returns whether the service is ready to handle requests
func (hc *HealthChecker) IsReady() bool {
hc.readyMu.RLock()
defer hc.readyMu.RUnlock()
return hc.ready
}
// GetStatus returns the current health status
func (hc *HealthChecker) GetStatus() HealthStatus {
uptime := time.Since(hc.startTime)
deps := make(map[string]string)
if hc.redisHealthy {
deps["redis"] = "healthy"
} else {
deps["redis"] = "unhealthy"
}
if hc.ipfsClient != nil {
if hc.ipfsHealthy {
deps["ipfs"] = "healthy"
} else {
deps["ipfs"] = "unhealthy"
}
} else {
deps["ipfs"] = "not_configured"
}
status := "healthy"
if !hc.ready {
status = "unhealthy"
}
return HealthStatus{
Status: status,
Timestamp: time.Now().UTC().Format(time.RFC3339),
Uptime: uptime.String(),
Dependencies: deps,
}
}
// Global health checker instance
var healthChecker *HealthChecker
// InitHealthChecker initializes the global health checker
func InitHealthChecker(redisClient *asynq.Client, ipfsClient ipfs.IPFSClient) {
healthChecker = NewHealthChecker(redisClient, ipfsClient)
}
// HealthCheckHandler returns health status (liveness probe)
func HealthCheckHandler(c echo.Context) error {
if healthChecker == nil {
return c.JSON(http.StatusOK, map[string]string{"status": "starting"})
}
status := healthChecker.GetStatus()
if status.Status == "healthy" {
return c.JSON(http.StatusOK, status)
}
return c.JSON(http.StatusServiceUnavailable, status)
}
// ReadinessHandler returns readiness status (readiness probe)
func ReadinessHandler(c echo.Context) error {
if healthChecker == nil || !healthChecker.IsReady() {
return c.JSON(http.StatusServiceUnavailable, map[string]string{
"ready": "false",
"reason": "service not ready",
})
}
return c.JSON(http.StatusOK, map[string]string{
"ready": "true",
})
}
+422
View File
@@ -0,0 +1,422 @@
package handlers
import (
"crypto/rand"
"encoding/base64"
"fmt"
"net/url"
"strings"
"sync"
"time"
)
const (
// ClientTypePublic represents a public OAuth2 client
ClientTypePublic = "public"
// ClientTypeConfidential represents a confidential OAuth2 client
ClientTypeConfidential = "confidential"
)
// ClientRegistry manages OAuth2 client registrations
type ClientRegistry struct {
mu sync.RWMutex
clients map[string]*OAuth2Client
}
// NewClientRegistry creates a new client registry
func NewClientRegistry() *ClientRegistry {
registry := &ClientRegistry{
clients: make(map[string]*OAuth2Client),
}
// Initialize with default clients for development
registry.initializeDefaultClients()
return registry
}
// initializeDefaultClients adds default clients for development/testing
func (r *ClientRegistry) initializeDefaultClients() {
// Development public client (e.g., SPA)
_ = r.RegisterClient(&OAuth2Client{
ClientID: "sonr-web-app",
ClientType: ClientTypePublic,
RedirectURIs: []string{
"http://localhost:3000/callback",
"http://localhost:3001/callback",
},
AllowedScopes: []string{
"openid",
"profile",
"vault:read",
"vault:write",
"service:manage",
},
AllowedGrants: []string{"authorization_code", "refresh_token"},
TokenLifetime: time.Hour,
RequirePKCE: true,
TrustedClient: true,
RequiresConsent: false,
Metadata: map[string]string{
"name": "Sonr Web Application",
"description": "Official Sonr web application",
"logo_uri": "https://sonr.io/logo.png",
"client_uri": "https://app.sonr.io",
},
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
})
// Development confidential client (e.g., backend service)
_ = r.RegisterClient(&OAuth2Client{
ClientID: "sonr-backend-service",
ClientSecret: "development-secret-change-in-production",
ClientType: ClientTypeConfidential,
RedirectURIs: []string{"http://localhost:8081/callback"},
AllowedScopes: []string{"openid", "profile", "vault:admin", "service:manage"},
AllowedGrants: []string{"authorization_code", "refresh_token", "client_credentials"},
TokenLifetime: time.Hour * 2,
RequirePKCE: false,
TrustedClient: true,
RequiresConsent: false,
Metadata: map[string]string{
"name": "Sonr Backend Service",
"description": "Backend service for Sonr ecosystem",
},
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
})
// Example third-party client
_ = r.RegisterClient(&OAuth2Client{
ClientID: "example-third-party",
ClientSecret: "third-party-secret",
ClientType: ClientTypeConfidential,
RedirectURIs: []string{"https://example.com/oauth/callback"},
AllowedScopes: []string{"openid", "profile", "vault:read"},
AllowedGrants: []string{"authorization_code", "refresh_token"},
TokenLifetime: time.Hour,
RequirePKCE: true,
TrustedClient: false,
RequiresConsent: true,
Metadata: map[string]string{
"name": "Example Third Party App",
"description": "Example integration partner",
"logo_uri": "https://example.com/logo.png",
"client_uri": "https://example.com",
"policy_uri": "https://example.com/privacy",
"tos_uri": "https://example.com/terms",
},
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
})
}
// RegisterClient registers a new OAuth2 client
func (r *ClientRegistry) RegisterClient(client *OAuth2Client) error {
r.mu.Lock()
defer r.mu.Unlock()
// Validate client
if err := r.validateClient(client); err != nil {
return err
}
// Generate client ID if not provided
if client.ClientID == "" {
client.ClientID = generateOAuth2ClientID()
}
// Generate client secret for confidential clients
if client.ClientType == ClientTypeConfidential && client.ClientSecret == "" {
client.ClientSecret = generateClientSecret()
}
// Set timestamps
if client.CreatedAt.IsZero() {
client.CreatedAt = time.Now()
}
client.UpdatedAt = time.Now()
// Store client
r.clients[client.ClientID] = client
return nil
}
// GetClient retrieves a client by ID
func (r *ClientRegistry) GetClient(clientID string) (*OAuth2Client, error) {
r.mu.RLock()
defer r.mu.RUnlock()
client, exists := r.clients[clientID]
if !exists {
return nil, fmt.Errorf("client not found: %s", clientID)
}
return client, nil
}
// UpdateClient updates an existing client
func (r *ClientRegistry) UpdateClient(clientID string, updates *OAuth2Client) error {
r.mu.Lock()
defer r.mu.Unlock()
client, exists := r.clients[clientID]
if !exists {
return fmt.Errorf("client not found: %s", clientID)
}
// Update allowed fields
if len(updates.RedirectURIs) > 0 {
client.RedirectURIs = updates.RedirectURIs
}
if len(updates.AllowedScopes) > 0 {
client.AllowedScopes = updates.AllowedScopes
}
if len(updates.AllowedGrants) > 0 {
client.AllowedGrants = updates.AllowedGrants
}
if updates.TokenLifetime > 0 {
client.TokenLifetime = updates.TokenLifetime
}
if updates.Metadata != nil {
client.Metadata = updates.Metadata
}
client.RequirePKCE = updates.RequirePKCE
client.RequiresConsent = updates.RequiresConsent
client.UpdatedAt = time.Now()
return nil
}
// DeleteClient removes a client from the registry
func (r *ClientRegistry) DeleteClient(clientID string) error {
r.mu.Lock()
defer r.mu.Unlock()
if _, exists := r.clients[clientID]; !exists {
return fmt.Errorf("client not found: %s", clientID)
}
delete(r.clients, clientID)
return nil
}
// ListClients returns all registered clients
func (r *ClientRegistry) ListClients() []*OAuth2Client {
r.mu.RLock()
defer r.mu.RUnlock()
clients := make([]*OAuth2Client, 0, len(r.clients))
for _, client := range r.clients {
clients = append(clients, client)
}
return clients
}
// ValidateRedirectURI checks if a redirect URI is valid for the client
func (c *OAuth2Client) ValidateRedirectURI(redirectURI string) bool {
if redirectURI == "" {
return false
}
// Parse the redirect URI
parsedURI, err := url.Parse(redirectURI)
if err != nil {
return false
}
// Check against registered redirect URIs
for _, registeredURI := range c.RedirectURIs {
registeredParsed, err := url.Parse(registeredURI)
if err != nil {
continue
}
// For public clients, allow localhost with any port
if c.ClientType == ClientTypePublic && registeredParsed.Hostname() == "localhost" &&
parsedURI.Hostname() == "localhost" {
if registeredParsed.Path == parsedURI.Path {
return true
}
}
// Exact match for other cases
if registeredURI == redirectURI {
return true
}
// Allow subdomain matching for trusted clients
if c.TrustedClient && matchesWithSubdomain(registeredParsed, parsedURI) {
return true
}
}
return false
}
// ValidateScopes checks if the requested scopes are allowed for the client
func (c *OAuth2Client) ValidateScopes(requestedScopes []string) bool {
if len(requestedScopes) == 0 {
return true // No scopes requested is valid
}
for _, scope := range requestedScopes {
if !c.hasScope(scope) {
return false
}
}
return true
}
// hasScope checks if a client has a specific scope
func (c *OAuth2Client) hasScope(scope string) bool {
for _, allowedScope := range c.AllowedScopes {
if allowedScope == scope {
return true
}
// Check for hierarchical scopes (e.g., vault:admin includes vault:read)
if isHierarchicalScope(allowedScope, scope) {
return true
}
}
return false
}
// HasGrantType checks if a client supports a specific grant type
func (c *OAuth2Client) HasGrantType(grantType string) bool {
for _, allowed := range c.AllowedGrants {
if allowed == grantType {
return true
}
}
return false
}
// validateClient validates client configuration
func (r *ClientRegistry) validateClient(client *OAuth2Client) error {
// Validate client type
if client.ClientType != ClientTypePublic && client.ClientType != ClientTypeConfidential {
return fmt.Errorf("invalid client type: %s", client.ClientType)
}
// Validate redirect URIs
if len(client.RedirectURIs) == 0 {
return fmt.Errorf("at least one redirect URI is required")
}
for _, uri := range client.RedirectURIs {
if _, err := url.Parse(uri); err != nil {
return fmt.Errorf("invalid redirect URI: %s", uri)
}
}
// Validate grant types
if len(client.AllowedGrants) == 0 {
client.AllowedGrants = []string{"authorization_code"}
}
validGrants := map[string]bool{
"authorization_code": true,
"implicit": true,
"refresh_token": true,
"client_credentials": true,
"password": true,
"urn:ietf:params:oauth:grant-type:device_code": true,
}
for _, grant := range client.AllowedGrants {
if !validGrants[grant] {
return fmt.Errorf("invalid grant type: %s", grant)
}
}
// Client credentials grant requires confidential client
if contains(client.AllowedGrants, "client_credentials") &&
client.ClientType != ClientTypeConfidential {
return fmt.Errorf("client_credentials grant requires confidential client")
}
// Public clients should use PKCE
if client.ClientType == ClientTypePublic && !client.RequirePKCE {
// Log warning but don't fail
fmt.Printf("Warning: Public client %s should use PKCE\n", client.ClientID)
}
return nil
}
// Helper functions
func generateOAuth2ClientID() string {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return fmt.Sprintf("client_%s", base64.RawURLEncoding.EncodeToString(bytes))
}
func generateClientSecret() string {
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
func matchesWithSubdomain(registered, requested *url.URL) bool {
if registered.Scheme != requested.Scheme {
return false
}
if registered.Path != requested.Path {
return false
}
// Check if requested hostname is a subdomain of registered
registeredHost := registered.Hostname()
requestedHost := requested.Hostname()
if registeredHost == requestedHost {
return true
}
// Check subdomain match (e.g., *.example.com matches sub.example.com)
if strings.HasPrefix(registeredHost, "*.") {
domain := strings.TrimPrefix(registeredHost, "*.")
return strings.HasSuffix(requestedHost, domain)
}
return false
}
func isHierarchicalScope(allowed, requested string) bool {
// Define scope hierarchy
hierarchy := map[string][]string{
"vault:admin": {"vault:write", "vault:read", "vault:sign"},
"vault:write": {"vault:read"},
"service:manage": {"service:read", "service:write"},
"did:write": {"did:read"},
}
childScopes, exists := hierarchy[allowed]
if !exists {
return false
}
for _, child := range childScopes {
if child == requested {
return true
}
// Recursive check for nested hierarchies
if isHierarchicalScope(child, requested) {
return true
}
}
return false
}
+323
View File
@@ -0,0 +1,323 @@
package handlers
import (
"encoding/json"
"fmt"
"time"
"github.com/sonr-io/sonr/crypto/ucan"
)
// UCANDelegator handles UCAN token generation for OAuth flows
type UCANDelegator struct {
scopeMapper *ScopeMapper
signer UCANSigner
}
// UCANSigner interface for signing UCAN tokens
type UCANSigner interface {
Sign(token *ucan.Token) (string, error)
GetIssuerDID() string
}
// NewUCANDelegator creates a new UCAN delegator
func NewUCANDelegator(signer UCANSigner) *UCANDelegator {
if signer == nil {
// Use blockchain signer by default
signer, _ = NewBlockchainUCANSigner(nil, "did:sonr:oauth-provider")
}
return &UCANDelegator{
scopeMapper: NewScopeMapper(),
signer: signer,
}
}
// CreateDelegation creates a UCAN token for user-to-client delegation
func (d *UCANDelegator) CreateDelegation(
userDID string,
clientID string,
scopes []string,
expiresAt time.Time,
) (*ucan.Token, error) {
// Build resource context for the user
resourceContext := d.buildResourceContext(userDID)
// Map OAuth scopes to UCAN attenuations
attenuations := d.scopeMapper.MapToUCAN(scopes, userDID, clientID, resourceContext)
if len(attenuations) == 0 {
return nil, fmt.Errorf("no valid attenuations for scopes: %v", scopes)
}
// Create UCAN token
token := &ucan.Token{
Issuer: userDID,
Audience: clientID,
ExpiresAt: expiresAt.Unix(),
NotBefore: time.Now().Unix(),
Attenuations: attenuations,
Facts: []ucan.Fact{
{
Data: d.createOAuthFact(scopes, "user_delegation"),
},
},
}
// Sign the token
signedToken, err := d.signToken(token)
if err != nil {
return nil, fmt.Errorf("failed to sign UCAN token: %w", err)
}
token.Raw = signedToken
return token, nil
}
// CreateServiceDelegation creates a UCAN token for service-to-service delegation
func (d *UCANDelegator) CreateServiceDelegation(
clientID string,
scopes []string,
expiresAt time.Time,
) (*ucan.Token, error) {
// Service delegations use the OAuth provider as issuer
issuerDID := d.signer.GetIssuerDID()
// Build resource context for service
resourceContext := map[string]string{
"service_id": clientID,
"type": "service",
}
// Map OAuth scopes to UCAN attenuations
attenuations := d.scopeMapper.MapToUCAN(scopes, issuerDID, clientID, resourceContext)
if len(attenuations) == 0 {
return nil, fmt.Errorf("no valid attenuations for scopes: %v", scopes)
}
// Create UCAN token
token := &ucan.Token{
Issuer: issuerDID,
Audience: clientID,
ExpiresAt: expiresAt.Unix(),
NotBefore: time.Now().Unix(),
Attenuations: attenuations,
Facts: []ucan.Fact{
{
Data: d.createOAuthFact(scopes, "service_delegation"),
},
},
}
// Sign the token
signedToken, err := d.signToken(token)
if err != nil {
return nil, fmt.Errorf("failed to sign UCAN token: %w", err)
}
token.Raw = signedToken
return token, nil
}
// CreateDelegationChain creates a chain of UCAN tokens for complex delegations
func (d *UCANDelegator) CreateDelegationChain(
userDID string,
intermediaries []string,
finalAudience string,
scopes []string,
expiresAt time.Time,
) ([]*ucan.Token, error) {
chain := make([]*ucan.Token, 0, len(intermediaries)+1)
currentIssuer := userDID
proofs := []ucan.Proof{}
// Create delegation for each intermediary
for _, intermediary := range intermediaries {
token, err := d.createIntermediateDelegation(
currentIssuer,
intermediary,
scopes,
expiresAt,
proofs,
)
if err != nil {
return nil, fmt.Errorf("failed to create intermediate delegation: %w", err)
}
chain = append(chain, token)
proofs = append(proofs, ucan.Proof(token.Raw))
currentIssuer = intermediary
}
// Create final delegation
finalToken, err := d.createFinalDelegation(
currentIssuer,
finalAudience,
scopes,
expiresAt,
proofs,
)
if err != nil {
return nil, fmt.Errorf("failed to create final delegation: %w", err)
}
chain = append(chain, finalToken)
return chain, nil
}
// RevokeDelegation revokes a UCAN delegation
func (d *UCANDelegator) RevokeDelegation(tokenID string) error {
// TODO: Implement delegation revocation
// This would typically involve:
// 1. Adding the token to a revocation list
// 2. Publishing revocation to a public ledger or database
// 3. Notifying relevant parties
return fmt.Errorf("delegation revocation not yet implemented")
}
// ValidateDelegation validates a UCAN token delegation
func (d *UCANDelegator) ValidateDelegation(token *ucan.Token, requiredScopes []string) error {
// Check expiration
if time.Now().Unix() > token.ExpiresAt {
return fmt.Errorf("token expired")
}
// Check not before
if token.NotBefore > 0 && time.Now().Unix() < token.NotBefore {
return fmt.Errorf("token not yet valid")
}
// Verify the token has required capabilities
for _, scope := range requiredScopes {
if !d.tokenGrantsScope(token, scope) {
return fmt.Errorf("token does not grant required scope: %s", scope)
}
}
// TODO: Verify signature
// TODO: Verify delegation chain if proofs exist
return nil
}
// Private helper methods
func (d *UCANDelegator) buildResourceContext(userDID string) map[string]string {
// TODO: Fetch actual resource information for the user
return map[string]string{
"vault_address": fmt.Sprintf("vault:%s", userDID),
"enclave_cid": fmt.Sprintf("cid:%s", userDID),
"dwn_id": fmt.Sprintf("dwn:%s", userDID),
"service_id": fmt.Sprintf("service:%s", userDID),
}
}
func (d *UCANDelegator) createOAuthFact(scopes []string, delegationType string) json.RawMessage {
fact := map[string]any{
"oauth_scopes": scopes,
"delegation_type": delegationType,
"issued_at": time.Now().Unix(),
"oauth_version": "2.0",
}
data, _ := json.Marshal(fact)
return json.RawMessage(data)
}
func (d *UCANDelegator) signToken(token *ucan.Token) (string, error) {
if d.signer == nil {
return "", fmt.Errorf("no signer configured")
}
return d.signer.Sign(token)
}
func (d *UCANDelegator) createIntermediateDelegation(
issuer, audience string,
scopes []string,
expiresAt time.Time,
proofs []ucan.Proof,
) (*ucan.Token, error) {
return d.createDelegationWithType(
issuer,
audience,
scopes,
expiresAt,
proofs,
"intermediate_delegation",
)
}
func (d *UCANDelegator) createFinalDelegation(
issuer, audience string,
scopes []string,
expiresAt time.Time,
proofs []ucan.Proof,
) (*ucan.Token, error) {
return d.createDelegationWithType(
issuer,
audience,
scopes,
expiresAt,
proofs,
"final_delegation",
)
}
// createDelegationWithType creates a delegation token with a specific type
func (d *UCANDelegator) createDelegationWithType(
issuer, audience string,
scopes []string,
expiresAt time.Time,
proofs []ucan.Proof,
delegationType string,
) (*ucan.Token, error) {
resourceContext := d.buildResourceContext(issuer)
attenuations := d.scopeMapper.MapToUCAN(scopes, issuer, audience, resourceContext)
token := &ucan.Token{
Issuer: issuer,
Audience: audience,
ExpiresAt: expiresAt.Unix(),
NotBefore: time.Now().Unix(),
Attenuations: attenuations,
Proofs: proofs,
Facts: []ucan.Fact{
{
Data: d.createOAuthFact(scopes, delegationType),
},
},
}
signedToken, err := d.signToken(token)
if err != nil {
return nil, err
}
token.Raw = signedToken
return token, nil
}
func (d *UCANDelegator) tokenGrantsScope(token *ucan.Token, scope string) bool {
// Get the scope definition
scopeDef, exists := d.scopeMapper.GetScope(scope)
if !exists {
return false
}
// Check if any attenuation grants the required capabilities
for _, attenuation := range token.Attenuations {
if attenuation.Capability.Grants(scopeDef.UCANActions) {
// Also check resource type matches
if d.resourceMatches(attenuation.Resource, scopeDef.ResourceType) {
return true
}
}
}
return false
}
func (d *UCANDelegator) resourceMatches(resource ucan.Resource, resourceType string) bool {
// Simple scheme matching for now
return resource.GetScheme() == resourceType || resource.GetScheme() == "*"
}
+923
View File
@@ -0,0 +1,923 @@
package handlers
import (
"crypto/rand"
"crypto/subtle"
"encoding/base64"
"fmt"
"net/http"
"strings"
"sync"
"time"
"github.com/labstack/echo/v4"
"github.com/sonr-io/sonr/crypto/ucan"
)
// OAuth2Provider extends OIDCProvider with full OAuth2 capabilities
type OAuth2Provider struct {
*OIDCProvider
clientRegistry *ClientRegistry
scopeMapper *ScopeMapper
ucanDelegator *UCANDelegator
authCodeStore *AuthCodeStore
accessTokenStore *AccessTokenStore
refreshTokenStore *RefreshTokenStore
consentStore *ConsentStore
config *OAuth2Config
}
// AuthCodeStore manages authorization codes
type AuthCodeStore struct {
mu sync.RWMutex
codes map[string]*OAuth2AuthorizationCode
}
// AccessTokenStore manages access tokens
type AccessTokenStore struct {
mu sync.RWMutex
tokens map[string]*OAuth2AccessToken
}
// RefreshTokenStore manages refresh tokens
type RefreshTokenStore struct {
mu sync.RWMutex
tokens map[string]*OAuth2RefreshToken
}
// ConsentStore manages user consent records
type ConsentStore struct {
mu sync.RWMutex
consents map[string]*UserConsent // key: userDID:clientID
}
// UserConsent represents stored user consent
type UserConsent struct {
UserDID string `json:"user_did"`
ClientID string `json:"client_id"`
ApprovedScopes []string `json:"approved_scopes"`
CreatedAt time.Time `json:"created_at"`
ExpiresAt time.Time `json:"expires_at"`
}
var oauth2Provider *OAuth2Provider
// InitializeOAuth2Provider initializes the OAuth2 provider
func InitializeOAuth2Provider() {
oauth2Provider = &OAuth2Provider{
OIDCProvider: oidcProvider,
clientRegistry: NewClientRegistry(),
scopeMapper: NewScopeMapper(),
ucanDelegator: NewUCANDelegator(nil),
authCodeStore: &AuthCodeStore{codes: make(map[string]*OAuth2AuthorizationCode)},
accessTokenStore: &AccessTokenStore{tokens: make(map[string]*OAuth2AccessToken)},
refreshTokenStore: &RefreshTokenStore{tokens: make(map[string]*OAuth2RefreshToken)},
consentStore: &ConsentStore{consents: make(map[string]*UserConsent)},
config: getDefaultOAuth2Config(),
}
// Start cleanup goroutine for expired tokens
go oauth2Provider.cleanupExpiredTokens()
}
// GetOAuth2Discovery returns OAuth2 discovery configuration
func GetOAuth2Discovery(c echo.Context) error {
if oauth2Provider == nil {
InitializeOAuth2Provider()
}
return c.JSON(http.StatusOK, oauth2Provider.config)
}
// HandleOAuth2Authorize handles OAuth2 authorization requests
func HandleOAuth2Authorize(c echo.Context) error {
if oauth2Provider == nil {
InitializeOAuth2Provider()
}
req := &OAuth2AuthorizationRequest{
ResponseType: c.QueryParam("response_type"),
ClientID: c.QueryParam("client_id"),
RedirectURI: c.QueryParam("redirect_uri"),
Scope: c.QueryParam("scope"),
State: c.QueryParam("state"),
CodeChallenge: c.QueryParam("code_challenge"),
CodeChallengeMethod: c.QueryParam("code_challenge_method"),
Nonce: c.QueryParam("nonce"),
Prompt: c.QueryParam("prompt"),
LoginHint: c.QueryParam("login_hint"),
}
// Validate client
client, err := oauth2Provider.clientRegistry.GetClient(req.ClientID)
if err != nil {
return oauth2Error(c, "invalid_client", "Unknown client", req.State)
}
// Validate redirect URI
if !client.ValidateRedirectURI(req.RedirectURI) {
return oauth2Error(c, "invalid_request", "Invalid redirect URI", req.State)
}
// Validate response type
if !isValidResponseType(req.ResponseType) {
return redirectError(
c,
req.RedirectURI,
"unsupported_response_type",
"Response type not supported",
req.State,
)
}
// Validate scopes
requestedScopes := parseScopes(req.Scope)
if !client.ValidateScopes(requestedScopes) {
return redirectError(
c,
req.RedirectURI,
"invalid_scope",
"Requested scope not allowed",
req.State,
)
}
// Validate PKCE for public clients
if client.ClientType == "public" && client.RequirePKCE {
if req.CodeChallenge == "" {
return redirectError(
c,
req.RedirectURI,
"invalid_request",
"PKCE required for public clients",
req.State,
)
}
if req.CodeChallengeMethod != PKCEMethodS256 {
return redirectError(
c,
req.RedirectURI,
"invalid_request",
"Only S256 PKCE method supported",
req.State,
)
}
}
// Check authentication
userDID := c.Get("user_did")
if userDID == nil {
// Store authorization request and redirect to authentication
sessionID := generateSessionID()
// TODO: Store auth request in session store
authURL := fmt.Sprintf(
"/auth/login?session_id=%s&return_to=%s",
sessionID,
c.Request().URL.String(),
)
return c.Redirect(http.StatusFound, authURL)
}
// Check consent
if client.RequiresConsent &&
!oauth2Provider.hasValidConsent(userDID.(string), req.ClientID, requestedScopes) {
// Render consent page
return renderOAuth2ConsentPage(c, req, client)
}
// Generate authorization code
code := generateSecureToken(32)
authCode := &OAuth2AuthorizationCode{
Code: code,
ClientID: req.ClientID,
UserDID: userDID.(string),
RedirectURI: req.RedirectURI,
Scopes: requestedScopes,
State: req.State,
Nonce: req.Nonce,
CodeChallenge: req.CodeChallenge,
CodeChallengeMethod: req.CodeChallengeMethod,
ExpiresAt: time.Now().Add(10 * time.Minute),
UCANContext: oauth2Provider.buildUCANContext(userDID.(string)),
}
// Store authorization code
oauth2Provider.authCodeStore.Store(authCode)
// Build redirect URL
redirectURL := buildAuthorizationRedirect(req.RedirectURI, code, req.State)
return c.Redirect(http.StatusFound, redirectURL)
}
// HandleOAuth2Token handles OAuth2 token requests
func HandleOAuth2Token(c echo.Context) error {
if oauth2Provider == nil {
InitializeOAuth2Provider()
}
var req OAuth2TokenRequest
if err := c.Bind(&req); err != nil {
return oauth2TokenError(c, "invalid_request", "Invalid token request")
}
// Authenticate client
client, err := oauth2Provider.authenticateClient(c, &req)
if err != nil {
return oauth2TokenError(c, "invalid_client", "Client authentication failed")
}
// Handle grant type
switch req.GrantType {
case "authorization_code":
return oauth2Provider.handleAuthorizationCodeGrant(c, client, &req)
case "refresh_token":
return oauth2Provider.handleRefreshTokenGrant(c, client, &req)
case "client_credentials":
return oauth2Provider.handleClientCredentialsGrant(c, client, &req)
default:
return oauth2TokenError(c, "unsupported_grant_type", "Grant type not supported")
}
}
// HandleOAuth2Introspection handles token introspection requests
func HandleOAuth2Introspection(c echo.Context) error {
if oauth2Provider == nil {
InitializeOAuth2Provider()
}
var req OAuth2IntrospectionRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, &OAuth2IntrospectionResponse{Active: false})
}
// Authenticate client
client, err := oauth2Provider.authenticateClient(c, &OAuth2TokenRequest{
ClientID: req.ClientID,
ClientSecret: req.ClientSecret,
})
if err != nil {
return c.JSON(http.StatusUnauthorized, &OAuth2IntrospectionResponse{Active: false})
}
// Introspect token
response := oauth2Provider.introspectToken(req.Token, req.TokenTypeHint, client)
return c.JSON(http.StatusOK, response)
}
// HandleOAuth2Revocation handles token revocation requests
func HandleOAuth2Revocation(c echo.Context) error {
if oauth2Provider == nil {
InitializeOAuth2Provider()
}
var req OAuth2RevocationRequest
if err := c.Bind(&req); err != nil {
return c.NoContent(http.StatusBadRequest)
}
// Authenticate client
client, err := oauth2Provider.authenticateClient(c, &OAuth2TokenRequest{
ClientID: req.ClientID,
ClientSecret: req.ClientSecret,
})
if err != nil {
return c.NoContent(http.StatusUnauthorized)
}
// Revoke token
oauth2Provider.revokeToken(req.Token, req.TokenTypeHint, client)
return c.NoContent(http.StatusOK)
}
// Private methods
func (p *OAuth2Provider) handleAuthorizationCodeGrant(
c echo.Context,
client *OAuth2Client,
req *OAuth2TokenRequest,
) error {
// Retrieve authorization code
authCode := p.authCodeStore.Exchange(req.Code)
if authCode == nil {
return oauth2TokenError(c, "invalid_grant", "Invalid authorization code")
}
// Validate code hasn't expired
if time.Now().After(authCode.ExpiresAt) {
return oauth2TokenError(c, "invalid_grant", "Authorization code expired")
}
// Validate client
if authCode.ClientID != client.ClientID {
return oauth2TokenError(c, "invalid_grant", "Code was issued to different client")
}
// Validate redirect URI
if authCode.RedirectURI != req.RedirectURI {
return oauth2TokenError(c, "invalid_grant", "Redirect URI mismatch")
}
// Validate PKCE if present
if authCode.CodeChallenge != "" {
if !p.validatePKCE(req.CodeVerifier, authCode.CodeChallenge, authCode.CodeChallengeMethod) {
return oauth2TokenError(c, "invalid_grant", "Invalid PKCE verifier")
}
}
// Create UCAN delegation
ucanToken, err := p.ucanDelegator.CreateDelegation(
authCode.UserDID,
client.ClientID,
authCode.Scopes,
time.Now().Add(time.Hour),
)
if err != nil {
return oauth2TokenError(c, "server_error", "Failed to create delegation")
}
// Generate tokens
accessToken := p.generateAccessToken(authCode, ucanToken)
refreshToken := p.generateRefreshToken(authCode)
// Store tokens
p.accessTokenStore.Store(accessToken)
p.refreshTokenStore.Store(refreshToken)
// Generate ID token if openid scope present
var idToken string
if contains(authCode.Scopes, "openid") {
idToken, _ = generateIDToken(authCode.UserDID, client.ClientID, authCode.Nonce)
}
// Return token response
response := &OAuth2TokenResponse{
AccessToken: accessToken.Token,
TokenType: "Bearer",
ExpiresIn: 3600,
RefreshToken: refreshToken.Token,
Scope: strings.Join(authCode.Scopes, " "),
IDToken: idToken,
UCANToken: ucanToken.Raw,
}
return c.JSON(http.StatusOK, response)
}
func (p *OAuth2Provider) handleRefreshTokenGrant(
c echo.Context,
client *OAuth2Client,
req *OAuth2TokenRequest,
) error {
// Retrieve refresh token
oldRefreshToken := p.refreshTokenStore.Get(req.RefreshToken)
if oldRefreshToken == nil {
return oauth2TokenError(c, "invalid_grant", "Invalid refresh token")
}
// Validate client
if oldRefreshToken.ClientID != client.ClientID {
return oauth2TokenError(c, "invalid_grant", "Token was issued to different client")
}
// Validate expiration
if time.Now().After(oldRefreshToken.ExpiresAt) {
return oauth2TokenError(c, "invalid_grant", "Refresh token expired")
}
// Rotate refresh token
p.refreshTokenStore.Revoke(req.RefreshToken)
// Create new UCAN delegation
ucanToken, err := p.ucanDelegator.CreateDelegation(
oldRefreshToken.UserDID,
client.ClientID,
oldRefreshToken.Scopes,
time.Now().Add(time.Hour),
)
if err != nil {
return oauth2TokenError(c, "server_error", "Failed to create delegation")
}
// Generate new tokens
newAccessToken := &OAuth2AccessToken{
Token: generateSecureToken(32),
UserDID: oldRefreshToken.UserDID,
ClientID: client.ClientID,
Scopes: oldRefreshToken.Scopes,
ExpiresAt: time.Now().Add(time.Hour),
IssuedAt: time.Now(),
UCANToken: ucanToken,
}
newRefreshToken := &OAuth2RefreshToken{
Token: generateSecureToken(32),
AccessToken: newAccessToken.Token,
ClientID: client.ClientID,
UserDID: oldRefreshToken.UserDID,
Scopes: oldRefreshToken.Scopes,
ExpiresAt: time.Now().Add(30 * 24 * time.Hour),
IssuedAt: time.Now(),
RotationCount: oldRefreshToken.RotationCount + 1,
}
// Store new tokens
p.accessTokenStore.Store(newAccessToken)
p.refreshTokenStore.Store(newRefreshToken)
// Return response
response := &OAuth2TokenResponse{
AccessToken: newAccessToken.Token,
TokenType: "Bearer",
ExpiresIn: 3600,
RefreshToken: newRefreshToken.Token,
Scope: strings.Join(newAccessToken.Scopes, " "),
UCANToken: ucanToken.Raw,
}
return c.JSON(http.StatusOK, response)
}
func (p *OAuth2Provider) handleClientCredentialsGrant(
c echo.Context,
client *OAuth2Client,
req *OAuth2TokenRequest,
) error {
// Client credentials grant is only for confidential clients
if client.ClientType != "confidential" {
return oauth2TokenError(
c,
"unauthorized_client",
"Client type not authorized for this grant",
)
}
// Parse requested scopes
scopes := parseScopes(req.Scope)
if !client.ValidateScopes(scopes) {
return oauth2TokenError(c, "invalid_scope", "Requested scope not allowed")
}
// Create service-to-service UCAN token
ucanToken, err := p.ucanDelegator.CreateServiceDelegation(
client.ClientID,
scopes,
time.Now().Add(time.Hour),
)
if err != nil {
return oauth2TokenError(c, "server_error", "Failed to create delegation")
}
// Generate access token
accessToken := &OAuth2AccessToken{
Token: generateSecureToken(32),
UserDID: "", // No user for client credentials
ClientID: client.ClientID,
Scopes: scopes,
ExpiresAt: time.Now().Add(time.Hour),
IssuedAt: time.Now(),
UCANToken: ucanToken,
TokenType: "client_credentials",
}
// Store token
p.accessTokenStore.Store(accessToken)
// Return response
response := &OAuth2TokenResponse{
AccessToken: accessToken.Token,
TokenType: "Bearer",
ExpiresIn: 3600,
Scope: strings.Join(scopes, " "),
UCANToken: ucanToken.Raw,
}
return c.JSON(http.StatusOK, response)
}
func (p *OAuth2Provider) authenticateClient(
c echo.Context,
req *OAuth2TokenRequest,
) (*OAuth2Client, error) {
// Try Basic Auth first
if username, password, ok := c.Request().BasicAuth(); ok {
client, err := p.clientRegistry.GetClient(username)
if err != nil {
return nil, err
}
if client.ClientType == "confidential" &&
subtle.ConstantTimeCompare([]byte(client.ClientSecret), []byte(password)) == 1 {
return client, nil
}
return nil, fmt.Errorf("invalid client credentials")
}
// Try client_secret_post
if req.ClientID != "" && req.ClientSecret != "" {
client, err := p.clientRegistry.GetClient(req.ClientID)
if err != nil {
return nil, err
}
if client.ClientType == "confidential" &&
subtle.ConstantTimeCompare([]byte(client.ClientSecret), []byte(req.ClientSecret)) == 1 {
return client, nil
}
return nil, fmt.Errorf("invalid client credentials")
}
// Try client_assertion (JWT)
if req.ClientAssertion != "" &&
req.ClientAssertionType == "urn:ietf:params:oauth:client-assertion-type:jwt-bearer" {
// TODO: Implement JWT client assertion validation
return nil, fmt.Errorf("JWT client assertion not yet implemented")
}
// Public client (no authentication)
if req.ClientID != "" {
client, err := p.clientRegistry.GetClient(req.ClientID)
if err != nil {
return nil, err
}
if client.ClientType == "public" {
return client, nil
}
}
return nil, fmt.Errorf("client authentication required")
}
func (p *OAuth2Provider) validatePKCE(verifier, challenge, method string) bool {
if method == "" {
method = PKCEMethodPlain
}
computed := computePKCEChallenge(verifier, method)
return subtle.ConstantTimeCompare([]byte(computed), []byte(challenge)) == 1
}
func (p *OAuth2Provider) hasValidConsent(userDID, clientID string, scopes []string) bool {
p.consentStore.mu.RLock()
defer p.consentStore.mu.RUnlock()
key := fmt.Sprintf("%s:%s", userDID, clientID)
consent, exists := p.consentStore.consents[key]
if !exists {
return false
}
// Check expiration
if time.Now().After(consent.ExpiresAt) {
return false
}
// Check all requested scopes are approved
for _, scope := range scopes {
if !contains(consent.ApprovedScopes, scope) {
return false
}
}
return true
}
func (p *OAuth2Provider) buildUCANContext(userDID string) *UCANAuthContext {
// TODO: Fetch actual vault and DID document data
return &UCANAuthContext{
VaultAddress: fmt.Sprintf("vault_%s", userDID),
EnclaveDataCID: fmt.Sprintf("cid_%s", userDID),
Capabilities: []string{"read", "write", "sign"},
}
}
func (p *OAuth2Provider) generateAccessToken(
authCode *OAuth2AuthorizationCode,
ucanToken *ucan.Token,
) *OAuth2AccessToken {
return &OAuth2AccessToken{
Token: generateSecureToken(32),
UserDID: authCode.UserDID,
ClientID: authCode.ClientID,
Scopes: authCode.Scopes,
ExpiresAt: time.Now().Add(time.Hour),
IssuedAt: time.Now(),
UCANToken: ucanToken,
SessionID: generateSessionID(),
TokenType: "authorization_code",
}
}
func (p *OAuth2Provider) generateRefreshToken(
authCode *OAuth2AuthorizationCode,
) *OAuth2RefreshToken {
return &OAuth2RefreshToken{
Token: generateSecureToken(32),
ClientID: authCode.ClientID,
UserDID: authCode.UserDID,
Scopes: authCode.Scopes,
ExpiresAt: time.Now().Add(30 * 24 * time.Hour),
IssuedAt: time.Now(),
RotationCount: 0,
}
}
func (p *OAuth2Provider) introspectToken(
token, tokenTypeHint string,
client *OAuth2Client,
) *OAuth2IntrospectionResponse {
// Try access token first
if accessToken := p.accessTokenStore.Get(token); accessToken != nil {
if accessToken.ClientID != client.ClientID {
return &OAuth2IntrospectionResponse{Active: false}
}
return &OAuth2IntrospectionResponse{
Active: time.Now().Before(accessToken.ExpiresAt),
Scope: strings.Join(accessToken.Scopes, " "),
ClientID: accessToken.ClientID,
Username: accessToken.UserDID,
TokenType: "Bearer",
ExpiresAt: accessToken.ExpiresAt.Unix(),
IssuedAt: accessToken.IssuedAt.Unix(),
Subject: accessToken.UserDID,
UCANToken: accessToken.UCANToken.Raw,
}
}
// Try refresh token
if refreshToken := p.refreshTokenStore.Get(token); refreshToken != nil {
if refreshToken.ClientID != client.ClientID {
return &OAuth2IntrospectionResponse{Active: false}
}
return &OAuth2IntrospectionResponse{
Active: time.Now().Before(refreshToken.ExpiresAt),
Scope: strings.Join(refreshToken.Scopes, " "),
ClientID: refreshToken.ClientID,
Username: refreshToken.UserDID,
TokenType: "refresh_token",
ExpiresAt: refreshToken.ExpiresAt.Unix(),
IssuedAt: refreshToken.IssuedAt.Unix(),
Subject: refreshToken.UserDID,
}
}
return &OAuth2IntrospectionResponse{Active: false}
}
func (p *OAuth2Provider) revokeToken(token, tokenTypeHint string, client *OAuth2Client) {
// Try to revoke as access token
if p.accessTokenStore.Revoke(token) {
return
}
// Try to revoke as refresh token
p.refreshTokenStore.Revoke(token)
}
func (p *OAuth2Provider) cleanupExpiredTokens() {
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for range ticker.C {
// Cleanup expired authorization codes
p.authCodeStore.CleanupExpired()
// Cleanup expired access tokens
p.accessTokenStore.CleanupExpired()
// Cleanup expired refresh tokens
p.refreshTokenStore.CleanupExpired()
}
}
// Store methods for token stores
func (s *AuthCodeStore) Store(code *OAuth2AuthorizationCode) {
s.mu.Lock()
defer s.mu.Unlock()
s.codes[code.Code] = code
}
func (s *AuthCodeStore) Exchange(code string) *OAuth2AuthorizationCode {
s.mu.Lock()
defer s.mu.Unlock()
authCode, exists := s.codes[code]
if !exists || authCode.Used {
return nil
}
authCode.Used = true
return authCode
}
func (s *AuthCodeStore) CleanupExpired() {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
for code, authCode := range s.codes {
if now.After(authCode.ExpiresAt) {
delete(s.codes, code)
}
}
}
func (s *AccessTokenStore) Store(token *OAuth2AccessToken) {
s.mu.Lock()
defer s.mu.Unlock()
s.tokens[token.Token] = token
}
func (s *AccessTokenStore) Get(token string) *OAuth2AccessToken {
s.mu.RLock()
defer s.mu.RUnlock()
return s.tokens[token]
}
func (s *AccessTokenStore) Revoke(token string) bool {
s.mu.Lock()
defer s.mu.Unlock()
if _, exists := s.tokens[token]; exists {
delete(s.tokens, token)
return true
}
return false
}
func (s *AccessTokenStore) CleanupExpired() {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
for token, accessToken := range s.tokens {
if now.After(accessToken.ExpiresAt) {
delete(s.tokens, token)
}
}
}
func (s *RefreshTokenStore) Store(token *OAuth2RefreshToken) {
s.mu.Lock()
defer s.mu.Unlock()
s.tokens[token.Token] = token
}
func (s *RefreshTokenStore) Get(token string) *OAuth2RefreshToken {
s.mu.RLock()
defer s.mu.RUnlock()
return s.tokens[token]
}
func (s *RefreshTokenStore) Revoke(token string) bool {
s.mu.Lock()
defer s.mu.Unlock()
if _, exists := s.tokens[token]; exists {
delete(s.tokens, token)
return true
}
return false
}
func (s *RefreshTokenStore) CleanupExpired() {
s.mu.Lock()
defer s.mu.Unlock()
now := time.Now()
for token, refreshToken := range s.tokens {
if now.After(refreshToken.ExpiresAt) {
delete(s.tokens, token)
}
}
}
// Helper functions
func generateSecureToken(bytes int) string {
b := make([]byte, bytes)
if _, err := rand.Read(b); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(b)
}
func parseScopes(scope string) []string {
if scope == "" {
return []string{}
}
return strings.Split(scope, " ")
}
func contains(slice []string, item string) bool {
for _, s := range slice {
if s == item {
return true
}
}
return false
}
func isValidResponseType(responseType string) bool {
validTypes := []string{
"code",
"token",
"id_token",
"code id_token",
"code token",
"id_token token",
"code id_token token",
}
return contains(validTypes, responseType)
}
func oauth2Error(c echo.Context, error, description, state string) error {
return c.JSON(http.StatusBadRequest, &OAuth2ErrorResponse{
Error: error,
ErrorDescription: description,
State: state,
})
}
func oauth2TokenError(c echo.Context, error, description string) error {
return c.JSON(http.StatusBadRequest, &OAuth2ErrorResponse{
Error: error,
ErrorDescription: description,
})
}
func redirectError(c echo.Context, redirectURI, error, description, state string) error {
url := fmt.Sprintf("%s?error=%s&error_description=%s&state=%s",
redirectURI, error, description, state)
return c.Redirect(http.StatusFound, url)
}
func buildAuthorizationRedirect(redirectURI, code, state string) string {
if strings.Contains(redirectURI, "?") {
return fmt.Sprintf("%s&code=%s&state=%s", redirectURI, code, state)
}
return fmt.Sprintf("%s?code=%s&state=%s", redirectURI, code, state)
}
func renderOAuth2ConsentPage(
c echo.Context,
req *OAuth2AuthorizationRequest,
client *OAuth2Client,
) error {
// TODO: Render actual consent page
return c.JSON(http.StatusOK, map[string]any{
"client": client,
"scopes": parseScopes(req.Scope),
"state": req.State,
"client_id": req.ClientID,
})
}
func getDefaultOAuth2Config() *OAuth2Config {
baseURL := "https://localhost:8080"
return &OAuth2Config{
Issuer: baseURL,
AuthorizationEndpoint: baseURL + "/oauth2/authorize",
TokenEndpoint: baseURL + "/oauth2/token",
UserInfoEndpoint: baseURL + "/oauth2/userinfo",
JWKSEndpoint: baseURL + "/oauth2/jwks",
RegistrationEndpoint: baseURL + "/oauth2/register",
IntrospectionEndpoint: baseURL + "/oauth2/introspect",
RevocationEndpoint: baseURL + "/oauth2/revoke",
ScopesSupported: []string{
"openid", "profile", "email", "offline_access",
"vault:read", "vault:write", "vault:sign", "vault:admin",
"service:manage", "did:read", "did:write",
},
ResponseTypesSupported: []string{
"code", "token", "id_token",
"code id_token", "code token",
"id_token token", "code id_token token",
},
ResponseModesSupported: []string{
"query", "fragment", "form_post",
},
GrantTypesSupported: []string{
"authorization_code", "implicit", "refresh_token",
"client_credentials", "urn:ietf:params:oauth:grant-type:device_code",
},
SubjectTypesSupported: []string{
"public", "pairwise",
},
IDTokenSigningAlgValuesSupported: []string{
"ES256", "RS256", "HS256",
},
TokenEndpointAuthMethodsSupported: []string{
"client_secret_basic", "client_secret_post",
"client_secret_jwt", "private_key_jwt", "none",
},
ClaimsSupported: []string{
"sub", "iss", "aud", "exp", "iat", "auth_time", "nonce",
"name", "given_name", "family_name", "middle_name", "nickname",
"preferred_username", "profile", "picture", "website", "email",
"email_verified", "did", "vault_id", "ucan_capabilities",
},
CodeChallengeMethodsSupported: []string{
PKCEMethodS256, PKCEMethodPlain,
},
ServiceDocumentation: baseURL + "/docs/oauth2",
UILocalesSupported: []string{"en-US"},
UCANSupported: true,
}
}
+563
View File
@@ -0,0 +1,563 @@
package handlers
import (
"context"
"encoding/json"
"fmt"
"net/http"
"strings"
"time"
"github.com/sonr-io/sonr/crypto/ucan"
)
// RefreshTokenHandler handles OAuth2 refresh token flows with UCAN chains
type RefreshTokenHandler struct {
delegator *UCANDelegator
signer *BlockchainUCANSigner
tokenStore TokenStore
clientStore ClientStore
}
// RefreshTokenRequest represents an OAuth2 refresh token request
type RefreshTokenRequest struct {
GrantType string `json:"grant_type"`
RefreshToken string `json:"refresh_token"`
Scope string `json:"scope,omitempty"`
ClientID string `json:"client_id,omitempty"`
ClientSecret string `json:"client_secret,omitempty"`
}
// RefreshTokenResponse represents an OAuth2 refresh token response
type RefreshTokenResponse struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
RefreshToken string `json:"refresh_token,omitempty"`
Scope string `json:"scope,omitempty"`
UCANToken string `json:"ucan_token,omitempty"`
}
// UCANRefreshMetadata stores metadata for UCAN refresh chains
type UCANRefreshMetadata struct {
OriginalIssuer string `json:"original_issuer"`
DelegationChain []string `json:"delegation_chain"`
RefreshCount int `json:"refresh_count"`
MaxRefreshCount int `json:"max_refresh_count"`
CreatedAt time.Time `json:"created_at"`
LastRefreshedAt time.Time `json:"last_refreshed_at"`
AttenuationPath []Attenuation `json:"attenuation_path"`
}
// Attenuation represents scope reduction in the delegation chain
type Attenuation struct {
FromScopes []string `json:"from_scopes"`
ToScopes []string `json:"to_scopes"`
Timestamp time.Time `json:"timestamp"`
Reason string `json:"reason,omitempty"`
}
// NewRefreshTokenHandler creates a new refresh token handler
func NewRefreshTokenHandler(
delegator *UCANDelegator,
signer *BlockchainUCANSigner,
tokenStore TokenStore,
clientStore ClientStore,
) *RefreshTokenHandler {
return &RefreshTokenHandler{
delegator: delegator,
signer: signer,
tokenStore: tokenStore,
clientStore: clientStore,
}
}
// HandleRefreshToken handles OAuth2 refresh token requests
func (h *RefreshTokenHandler) HandleRefreshToken(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Parse request (handle both JSON and form-encoded)
var req RefreshTokenRequest
contentType := r.Header.Get("Content-Type")
if strings.Contains(contentType, "application/json") {
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
h.sendError(w, "invalid_request", "Failed to parse JSON request body")
return
}
} else {
// Parse form data
if err := r.ParseForm(); err != nil {
h.sendError(w, "invalid_request", "Failed to parse form data")
return
}
req.GrantType = r.FormValue("grant_type")
req.RefreshToken = r.FormValue("refresh_token")
req.Scope = r.FormValue("scope")
req.ClientID = r.FormValue("client_id")
req.ClientSecret = r.FormValue("client_secret")
}
// Validate grant type
if req.GrantType != "refresh_token" {
h.sendError(w, "unsupported_grant_type", "Only refresh_token grant type is supported")
return
}
// Validate refresh token
if req.RefreshToken == "" {
h.sendError(w, "invalid_request", "Missing refresh_token parameter")
return
}
// Authenticate client
clientID, clientSecret := h.extractClientCredentials(r, &req)
ctx := r.Context()
if err := h.clientStore.ValidateClientCredentials(ctx, clientID, clientSecret); err != nil {
h.sendError(w, "invalid_client", "Client authentication failed")
return
}
// Get client information
client, err := h.clientStore.GetClient(ctx, clientID)
if err != nil {
h.sendError(w, "invalid_client", "Client not found")
return
}
// Process refresh token
response, err := h.processRefreshToken(ctx, &req, client)
if err != nil {
h.sendError(w, "invalid_grant", err.Error())
return
}
// Send response
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
json.NewEncoder(w).Encode(response)
}
// processRefreshToken processes the refresh token and returns new tokens
func (h *RefreshTokenHandler) processRefreshToken(
ctx context.Context,
req *RefreshTokenRequest,
client *OAuth2Client,
) (*RefreshTokenResponse, error) {
// Retrieve stored refresh token
storedToken, err := h.tokenStore.GetToken(ctx, req.RefreshToken)
if err != nil {
return nil, fmt.Errorf("invalid refresh token")
}
// Validate token type
if storedToken.TokenType != "refresh_token" {
return nil, fmt.Errorf("token is not a refresh token")
}
// Validate client binding
if storedToken.ClientID != client.ClientID {
return nil, fmt.Errorf("refresh token was issued to a different client")
}
// Check if refresh token has expired
if time.Now().After(storedToken.ExpiresAt) {
return nil, fmt.Errorf("refresh token has expired")
}
// Get refresh metadata
metadata, err := h.getRefreshMetadata(ctx, req.RefreshToken)
if err != nil {
// Initialize metadata for first refresh
metadata = &UCANRefreshMetadata{
OriginalIssuer: storedToken.UserDID,
DelegationChain: []string{},
RefreshCount: 0,
MaxRefreshCount: 10, // Default max refresh count
CreatedAt: time.Now(),
AttenuationPath: []Attenuation{},
}
}
// Check refresh count limit
if metadata.RefreshCount >= metadata.MaxRefreshCount {
return nil, fmt.Errorf("refresh token has reached maximum refresh count")
}
// Parse requested scopes
requestedScopes := storedToken.Scopes
if req.Scope != "" {
requestedScopes = strings.Split(req.Scope, " ")
// Validate scope reduction (attenuate permissions)
if !h.validateScopeAttenuation(requestedScopes, storedToken.Scopes) {
return nil, fmt.Errorf("requested scopes exceed refresh token scopes")
}
// Record attenuation
if !h.scopesEqual(requestedScopes, storedToken.Scopes) {
metadata.AttenuationPath = append(metadata.AttenuationPath, Attenuation{
FromScopes: storedToken.Scopes,
ToScopes: requestedScopes,
Timestamp: time.Now(),
Reason: "Client requested scope reduction",
})
}
}
// Create new UCAN token with delegation chain
newUCANToken, err := h.createRefreshedUCANToken(
ctx,
metadata,
storedToken,
client,
requestedScopes,
)
if err != nil {
return nil, fmt.Errorf("failed to create refreshed UCAN token: %w", err)
}
// Generate new access token
accessTokenID := h.generateTokenID()
accessToken := &StoredToken{
TokenID: accessTokenID,
TokenType: "access_token",
AccessToken: accessTokenID,
ExpiresAt: time.Now().Add(time.Hour),
Scopes: requestedScopes,
ClientID: client.ClientID,
UserDID: storedToken.UserDID,
UCANToken: newUCANToken,
}
if err := h.tokenStore.StoreToken(ctx, accessToken); err != nil {
return nil, fmt.Errorf("failed to store new access token: %w", err)
}
// Update refresh metadata
metadata.RefreshCount++
metadata.LastRefreshedAt = time.Now()
metadata.DelegationChain = append(metadata.DelegationChain, newUCANToken)
// Optionally rotate refresh token
newRefreshTokenID := ""
if h.shouldRotateRefreshToken(metadata) {
newRefreshTokenID = h.generateTokenID()
newRefreshToken := &StoredToken{
TokenID: newRefreshTokenID,
TokenType: "refresh_token",
RefreshToken: newRefreshTokenID,
ExpiresAt: time.Now().Add(30 * 24 * time.Hour), // 30 days
Scopes: requestedScopes,
ClientID: client.ClientID,
UserDID: storedToken.UserDID,
}
if err := h.tokenStore.StoreToken(ctx, newRefreshToken); err != nil {
// Non-fatal, continue with existing refresh token
newRefreshTokenID = ""
} else {
// Revoke old refresh token
h.tokenStore.RevokeToken(ctx, req.RefreshToken)
// Store metadata for new refresh token
h.storeRefreshMetadata(ctx, newRefreshTokenID, metadata)
}
} else {
// Update metadata for existing refresh token
h.storeRefreshMetadata(ctx, req.RefreshToken, metadata)
}
// Build response
response := &RefreshTokenResponse{
AccessToken: accessTokenID,
TokenType: "Bearer",
ExpiresIn: 3600,
Scope: strings.Join(requestedScopes, " "),
UCANToken: newUCANToken,
}
if newRefreshTokenID != "" {
response.RefreshToken = newRefreshTokenID
}
return response, nil
}
// createRefreshedUCANToken creates a new UCAN token with proper delegation chain
func (h *RefreshTokenHandler) createRefreshedUCANToken(
ctx context.Context,
metadata *UCANRefreshMetadata,
storedToken *StoredToken,
client *OAuth2Client,
scopes []string,
) (string, error) {
// Build proof chain from previous delegations
proofs := make([]ucan.Proof, 0, len(metadata.DelegationChain))
for _, tokenStr := range metadata.DelegationChain {
proofs = append(proofs, ucan.Proof(tokenStr))
}
// Add original token as proof if exists
if storedToken.UCANToken != "" {
proofs = append([]ucan.Proof{ucan.Proof(storedToken.UCANToken)}, proofs...)
}
// Determine issuer and audience
issuer := metadata.OriginalIssuer
if issuer == "" {
issuer = storedToken.UserDID
}
audience := client.ClientID
if did, ok := client.Metadata["client_did"]; ok {
audience = did
}
// Create resource context with refresh metadata
resourceContext := map[string]string{
"refresh_count": fmt.Sprintf("%d", metadata.RefreshCount),
"original_issuer": metadata.OriginalIssuer,
"delegation_type": "refresh_token",
"client_id": client.ClientID,
}
// Map OAuth scopes to UCAN attenuations
attenuations := h.delegator.scopeMapper.MapToUCAN(scopes, issuer, audience, resourceContext)
// Create UCAN token with delegation chain
ucanToken := &ucan.Token{
Issuer: issuer,
Audience: audience,
ExpiresAt: time.Now().Add(time.Hour).Unix(),
NotBefore: time.Now().Unix(),
Attenuations: attenuations,
Proofs: proofs,
Facts: []ucan.Fact{
{
Data: h.createRefreshFact(metadata, scopes),
},
},
}
// Sign the token
signedToken, err := h.signer.Sign(ucanToken)
if err != nil {
return "", fmt.Errorf("failed to sign UCAN token: %w", err)
}
// Validate the delegation chain
if len(metadata.DelegationChain) > 0 {
allTokens := append([]string{storedToken.UCANToken}, metadata.DelegationChain...)
allTokens = append(allTokens, signedToken)
if err := h.signer.ValidateDelegationChain(allTokens); err != nil {
return "", fmt.Errorf("invalid delegation chain: %w", err)
}
}
return signedToken, nil
}
// validateScopeAttenuation validates that requested scopes are properly attenuated
func (h *RefreshTokenHandler) validateScopeAttenuation(requested, allowed []string) bool {
// Build allowed scope map
allowedMap := make(map[string]bool)
for _, scope := range allowed {
allowedMap[scope] = true
}
// Check each requested scope
for _, scope := range requested {
if !allowedMap[scope] {
// Check if a parent scope allows this
found := false
for _, allowedScope := range allowed {
if h.delegator.scopeMapper.IsHierarchicalScope(allowedScope, scope) {
found = true
break
}
}
if !found {
return false
}
}
}
return true
}
// shouldRotateRefreshToken determines if refresh token should be rotated
func (h *RefreshTokenHandler) shouldRotateRefreshToken(metadata *UCANRefreshMetadata) bool {
// Rotate on every use for maximum security
// Could be configured based on policy
return true
}
// extractClientCredentials extracts client credentials from request
func (h *RefreshTokenHandler) extractClientCredentials(
r *http.Request,
req *RefreshTokenRequest,
) (string, string) {
// Try Basic Auth first
if clientID, clientSecret, ok := r.BasicAuth(); ok {
return clientID, clientSecret
}
// Fall back to request body
return req.ClientID, req.ClientSecret
}
// getRefreshMetadata retrieves refresh metadata from storage
func (h *RefreshTokenHandler) getRefreshMetadata(
ctx context.Context,
refreshTokenID string,
) (*UCANRefreshMetadata, error) {
// In production, this would retrieve from persistent storage
// For now, return error to initialize new metadata
return nil, fmt.Errorf("metadata not found")
}
// storeRefreshMetadata stores refresh metadata
func (h *RefreshTokenHandler) storeRefreshMetadata(
ctx context.Context,
refreshTokenID string,
metadata *UCANRefreshMetadata,
) error {
// In production, this would persist to storage
// For now, just return success
return nil
}
// createRefreshFact creates a fact for refresh token
func (h *RefreshTokenHandler) createRefreshFact(
metadata *UCANRefreshMetadata,
scopes []string,
) json.RawMessage {
fact := map[string]any{
"type": "refresh_token",
"refresh_count": metadata.RefreshCount,
"original_issuer": metadata.OriginalIssuer,
"scopes": scopes,
"refreshed_at": time.Now().Unix(),
"delegation_length": len(metadata.DelegationChain),
}
// Add attenuation info if present
if len(metadata.AttenuationPath) > 0 {
fact["attenuations"] = len(metadata.AttenuationPath)
lastAttenuation := metadata.AttenuationPath[len(metadata.AttenuationPath)-1]
fact["last_attenuation"] = map[string]any{
"from": strings.Join(lastAttenuation.FromScopes, " "),
"to": strings.Join(lastAttenuation.ToScopes, " "),
"at": lastAttenuation.Timestamp.Unix(),
}
}
data, _ := json.Marshal(fact)
return json.RawMessage(data)
}
// scopesEqual checks if two scope slices are equal
func (h *RefreshTokenHandler) scopesEqual(a, b []string) bool {
if len(a) != len(b) {
return false
}
aMap := make(map[string]bool)
for _, scope := range a {
aMap[scope] = true
}
for _, scope := range b {
if !aMap[scope] {
return false
}
}
return true
}
// generateTokenID generates a unique token identifier
func (h *RefreshTokenHandler) generateTokenID() string {
// In production, use a proper UUID or secure random generator
return fmt.Sprintf("tok_%d_%s", time.Now().UnixNano(), h.randomString(16))
}
// randomString generates a random string
func (h *RefreshTokenHandler) randomString(length int) string {
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
result := make([]byte, length)
for i := range result {
result[i] = charset[time.Now().UnixNano()%int64(len(charset))]
}
return string(result)
}
// sendError sends an OAuth error response
func (h *RefreshTokenHandler) sendError(w http.ResponseWriter, errorCode, errorDescription string) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
w.WriteHeader(http.StatusBadRequest)
response := map[string]string{
"error": errorCode,
"error_description": errorDescription,
}
json.NewEncoder(w).Encode(response)
}
// HandleUCANRefresh handles UCAN-specific refresh requests
func (h *RefreshTokenHandler) HandleUCANRefresh(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Parse UCAN token from Authorization header
authHeader := r.Header.Get("Authorization")
if !strings.HasPrefix(authHeader, "Bearer ") {
h.sendError(w, "invalid_request", "Missing or invalid Authorization header")
return
}
tokenString := strings.TrimPrefix(authHeader, "Bearer ")
// Verify the UCAN token
ucanToken, err := h.signer.VerifySignature(tokenString)
if err != nil {
h.sendError(w, "invalid_grant", "Invalid UCAN token")
return
}
// Check if token can be refreshed (not expired beyond grace period)
gracePeriod := int64(300) // 5 minutes grace period
if time.Now().Unix() > ucanToken.ExpiresAt+gracePeriod {
h.sendError(w, "invalid_grant", "Token expired beyond grace period")
return
}
// Create refreshed token with extended expiration
newToken, err := h.signer.RefreshToken(tokenString, time.Hour)
if err != nil {
h.sendError(w, "server_error", "Failed to refresh token")
return
}
// Send response
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
json.NewEncoder(w).Encode(map[string]any{
"ucan_token": newToken,
"token_type": "UCAN",
"expires_in": 3600,
})
}
+383
View File
@@ -0,0 +1,383 @@
package handlers
import (
"crypto/rand"
"encoding/base64"
"encoding/json"
"net/http"
"strings"
"time"
)
// DynamicClientRegistrationRequest represents a client registration request per RFC 7591
type DynamicClientRegistrationRequest struct {
ClientName string `json:"client_name"`
RedirectURIs []string `json:"redirect_uris"`
GrantTypes []string `json:"grant_types,omitempty"`
ResponseTypes []string `json:"response_types,omitempty"`
Scope string `json:"scope,omitempty"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
ApplicationType string `json:"application_type,omitempty"`
Contacts []string `json:"contacts,omitempty"`
LogoURI string `json:"logo_uri,omitempty"`
ClientURI string `json:"client_uri,omitempty"`
PolicyURI string `json:"policy_uri,omitempty"`
TosURI string `json:"tos_uri,omitempty"`
JwksURI string `json:"jwks_uri,omitempty"`
SoftwareID string `json:"software_id,omitempty"`
SoftwareVersion string `json:"software_version,omitempty"`
}
// DynamicClientRegistrationResponse represents the response for client registration
type DynamicClientRegistrationResponse struct {
ClientID string `json:"client_id"`
ClientSecret string `json:"client_secret,omitempty"`
ClientName string `json:"client_name"`
RedirectURIs []string `json:"redirect_uris"`
GrantTypes []string `json:"grant_types"`
ResponseTypes []string `json:"response_types"`
Scope string `json:"scope"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
ApplicationType string `json:"application_type"`
ClientIDIssuedAt int64 `json:"client_id_issued_at"`
ClientSecretExpiresAt int64 `json:"client_secret_expires_at,omitempty"`
LogoURI string `json:"logo_uri,omitempty"`
ClientURI string `json:"client_uri,omitempty"`
PolicyURI string `json:"policy_uri,omitempty"`
TosURI string `json:"tos_uri,omitempty"`
JwksURI string `json:"jwks_uri,omitempty"`
}
// HandleDynamicClientRegistration handles dynamic client registration per RFC 7591
func (s *OAuth2Provider) HandleDynamicClientRegistration(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Parse registration request
var req DynamicClientRegistrationRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
http.Error(w, "Invalid registration request", http.StatusBadRequest)
return
}
// Validate required fields
if req.ClientName == "" {
http.Error(w, "client_name is required", http.StatusBadRequest)
return
}
if len(req.RedirectURIs) == 0 {
http.Error(w, "redirect_uris is required", http.StatusBadRequest)
return
}
// Validate redirect URIs
for _, uri := range req.RedirectURIs {
if !isValidRedirectURI(uri) {
http.Error(w, "Invalid redirect URI: "+uri, http.StatusBadRequest)
return
}
}
// Set defaults if not provided
if len(req.GrantTypes) == 0 {
req.GrantTypes = []string{"authorization_code"}
}
if len(req.ResponseTypes) == 0 {
req.ResponseTypes = []string{"code"}
}
if req.TokenEndpointAuthMethod == "" {
// Default based on application type
if req.ApplicationType == "native" || req.ApplicationType == "browser" {
req.TokenEndpointAuthMethod = "none" // Public client
} else {
req.TokenEndpointAuthMethod = "client_secret_basic"
}
}
if req.ApplicationType == "" {
req.ApplicationType = "web"
}
// Validate grant types and response types
if !validateGrantTypes(req.GrantTypes) {
http.Error(w, "Invalid grant types", http.StatusBadRequest)
return
}
if !validateResponseTypes(req.ResponseTypes) {
http.Error(w, "Invalid response types", http.StatusBadRequest)
return
}
// Validate scopes if scope mapper is available
if req.Scope != "" && s.scopeMapper != nil {
scopes := strings.Split(req.Scope, " ")
for _, scope := range scopes {
// Check if scope is valid using the scope mapper
if _, exists := s.scopeMapper.GetScope(scope); !exists {
http.Error(w, "Invalid scope: "+scope, http.StatusBadRequest)
return
}
}
}
// Generate client credentials
clientID := generateDynamicClientID()
var clientSecret string
var clientSecretExpiresAt int64
// Only generate secret for confidential clients
if req.TokenEndpointAuthMethod != "none" {
clientSecret = generateDynamicClientSecret()
// Client secrets expire in 1 year by default
clientSecretExpiresAt = time.Now().Add(365 * 24 * time.Hour).Unix()
}
// Create OAuth2 client
client := &OAuth2Client{
ClientID: clientID,
ClientSecret: clientSecret,
RedirectURIs: req.RedirectURIs,
AllowedScopes: strings.Split(req.Scope, " "),
AllowedGrants: req.GrantTypes,
TokenLifetime: time.Hour, // Default 1 hour
RequirePKCE: false,
TrustedClient: false,
RequiresConsent: true,
Metadata: map[string]string{
"client_name": req.ClientName,
"application_type": req.ApplicationType,
"logo_uri": req.LogoURI,
"client_uri": req.ClientURI,
"policy_uri": req.PolicyURI,
"tos_uri": req.TosURI,
"jwks_uri": req.JwksURI,
},
CreatedAt: time.Now(),
UpdatedAt: time.Now(),
}
// Set client type based on auth method
if req.TokenEndpointAuthMethod == "none" {
client.ClientType = "public"
} else {
client.ClientType = "confidential"
}
// Determine if client requires PKCE
if req.ApplicationType == "native" || req.ApplicationType == "browser" {
client.RequirePKCE = true
}
// Store client in the registry
if s.clientRegistry != nil {
if err := s.clientRegistry.RegisterClient(client); err != nil {
http.Error(w, "Failed to register client", http.StatusInternalServerError)
return
}
}
// Build response
resp := DynamicClientRegistrationResponse{
ClientID: clientID,
ClientSecret: clientSecret,
ClientName: req.ClientName,
RedirectURIs: req.RedirectURIs,
GrantTypes: req.GrantTypes,
ResponseTypes: req.ResponseTypes,
Scope: req.Scope,
TokenEndpointAuthMethod: req.TokenEndpointAuthMethod,
ApplicationType: req.ApplicationType,
ClientIDIssuedAt: time.Now().Unix(),
ClientSecretExpiresAt: clientSecretExpiresAt,
LogoURI: req.LogoURI,
ClientURI: req.ClientURI,
PolicyURI: req.PolicyURI,
TosURI: req.TosURI,
JwksURI: req.JwksURI,
}
// Return registration response
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(http.StatusCreated)
json.NewEncoder(w).Encode(resp)
}
// HandleClientConfiguration handles client configuration retrieval
func (s *OAuth2Provider) HandleClientConfiguration(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Extract client ID from path or query
clientID := r.URL.Query().Get("client_id")
if clientID == "" {
// Try to extract from path (e.g., /register/{client_id})
parts := strings.Split(r.URL.Path, "/")
if len(parts) > 2 {
clientID = parts[len(parts)-1]
}
}
if clientID == "" {
http.Error(w, "client_id is required", http.StatusBadRequest)
return
}
// Validate access token for client management
authHeader := r.Header.Get("Authorization")
if !strings.HasPrefix(authHeader, "Bearer ") {
w.Header().Set("WWW-Authenticate", `Bearer realm="client_configuration"`)
http.Error(w, "Access token required", http.StatusUnauthorized)
return
}
token := strings.TrimPrefix(authHeader, "Bearer ")
// Validate token in the access token store
if s.accessTokenStore == nil {
http.Error(w, "Token store not available", http.StatusInternalServerError)
return
}
// Check token validity (simplified for now)
// In production, this should validate the token properly
if token == "" {
http.Error(w, "Invalid or expired access token", http.StatusUnauthorized)
return
}
// Get client from registry
if s.clientRegistry == nil {
http.Error(w, "Client registry not available", http.StatusInternalServerError)
return
}
client, err := s.clientRegistry.GetClient(clientID)
if err != nil {
http.Error(w, "Client not found", http.StatusNotFound)
return
}
// Build response
resp := DynamicClientRegistrationResponse{
ClientID: client.ClientID,
ClientName: client.Metadata["client_name"],
RedirectURIs: client.RedirectURIs,
GrantTypes: client.AllowedGrants,
ResponseTypes: []string{"code", "token"}, // Default response types
Scope: strings.Join(client.AllowedScopes, " "),
TokenEndpointAuthMethod: getTokenEndpointAuthMethod(client),
ApplicationType: client.Metadata["application_type"],
ClientIDIssuedAt: client.CreatedAt.Unix(),
LogoURI: client.Metadata["logo_uri"],
ClientURI: client.Metadata["client_uri"],
PolicyURI: client.Metadata["policy_uri"],
TosURI: client.Metadata["tos_uri"],
JwksURI: client.Metadata["jwks_uri"],
}
// Return client configuration
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
json.NewEncoder(w).Encode(resp)
}
// getTokenEndpointAuthMethod determines the auth method from client type
func getTokenEndpointAuthMethod(client *OAuth2Client) string {
if client.ClientType == "public" {
return "none"
}
return "client_secret_basic"
}
// Helper functions for dynamic registration
func generateDynamicClientID() string {
// Generate a random client ID for dynamic registration
b := make([]byte, 16)
_, _ = rand.Read(b)
return "dyn_client_" + base64.RawURLEncoding.EncodeToString(b)
}
func generateDynamicClientSecret() string {
// Generate a secure random secret for dynamic registration
b := make([]byte, 32)
_, _ = rand.Read(b)
return base64.RawURLEncoding.EncodeToString(b)
}
func isValidRedirectURI(uri string) bool {
// Basic validation - in production, this should be more comprehensive
if uri == "" {
return false
}
// Allow localhost for development
if strings.HasPrefix(uri, "http://localhost") || strings.HasPrefix(uri, "http://127.0.0.1") {
return true
}
// Require HTTPS for production URIs
if !strings.HasPrefix(uri, "https://") {
// Allow custom schemes for native apps
if strings.Contains(uri, "://") {
return true
}
return false
}
return true
}
func validateGrantTypes(grantTypes []string) bool {
validGrants := map[string]bool{
"authorization_code": true,
"implicit": true,
"refresh_token": true,
"client_credentials": true,
"password": true,
}
for _, grant := range grantTypes {
if !validGrants[grant] {
return false
}
}
return true
}
func validateResponseTypes(responseTypes []string) bool {
validTypes := map[string]bool{
"code": true,
"token": true,
"id_token": true,
}
for _, respType := range responseTypes {
// Handle composite types like "code id_token"
parts := strings.Split(respType, " ")
for _, part := range parts {
if !validTypes[part] {
return false
}
}
}
return true
}
func hasScope(scopes []string, requiredScope string) bool {
for _, scope := range scopes {
if scope == requiredScope {
return true
}
}
return false
}
+412
View File
@@ -0,0 +1,412 @@
package handlers
import (
"fmt"
"strings"
"github.com/sonr-io/sonr/crypto/ucan"
)
// ScopeMapper manages OAuth scope to UCAN capability mapping
type ScopeMapper struct {
scopeDefinitions map[string]*OAuth2ScopeDefinition
ucanTemplates map[string]*ucan.Attenuation
}
// NewScopeMapper creates a new scope mapper with standard scopes
func NewScopeMapper() *ScopeMapper {
mapper := &ScopeMapper{
scopeDefinitions: make(map[string]*OAuth2ScopeDefinition),
ucanTemplates: make(map[string]*ucan.Attenuation),
}
mapper.initializeStandardScopes()
return mapper
}
// initializeStandardScopes defines the standard OAuth scopes and their UCAN mappings
func (m *ScopeMapper) initializeStandardScopes() {
// OpenID Connect scopes
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "openid",
Description: "OpenID Connect authentication",
UCANActions: []string{"authenticate"},
ResourceType: "identity",
RequiresAuth: true,
Sensitive: false,
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "profile",
Description: "Access to user profile information",
UCANActions: []string{"read"},
ResourceType: "did",
RequiresAuth: true,
Sensitive: false,
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "email",
Description: "Access to user email",
UCANActions: []string{"read"},
ResourceType: "contact",
RequiresAuth: true,
Sensitive: true,
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "offline_access",
Description: "Maintain access when user is not present",
UCANActions: []string{"refresh"},
ResourceType: "session",
RequiresAuth: true,
Sensitive: true,
})
// Vault scopes
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "vault:read",
Description: "Read access to vault data",
UCANActions: []string{"read"},
ResourceType: "vault",
RequiresAuth: true,
Sensitive: false,
ParentScope: "",
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "vault:write",
Description: "Write access to vault data",
UCANActions: []string{"read", "write"},
ResourceType: "vault",
RequiresAuth: true,
Sensitive: true,
ParentScope: "vault:read",
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "vault:sign",
Description: "Signing operations with vault keys",
UCANActions: []string{"read", "sign"},
ResourceType: "vault",
RequiresAuth: true,
Sensitive: true,
ParentScope: "vault:read",
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "vault:admin",
Description: "Full administrative access to vault",
UCANActions: []string{"read", "write", "sign", "export", "import", "delete", "admin"},
ResourceType: "vault",
RequiresAuth: true,
Sensitive: true,
ParentScope: "vault:write",
ChildScopes: []string{"vault:read", "vault:write", "vault:sign"},
})
// Service scopes
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "service:read",
Description: "Read service information",
UCANActions: []string{"read"},
ResourceType: "service",
RequiresAuth: true,
Sensitive: false,
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "service:write",
Description: "Create and update services",
UCANActions: []string{"read", "write"},
ResourceType: "service",
RequiresAuth: true,
Sensitive: true,
ParentScope: "service:read",
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "service:manage",
Description: "Full service management capabilities",
UCANActions: []string{"read", "write", "delete", "admin"},
ResourceType: "service",
RequiresAuth: true,
Sensitive: true,
ParentScope: "service:write",
ChildScopes: []string{"service:read", "service:write"},
})
// DID scopes
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "did:read",
Description: "Read DID documents",
UCANActions: []string{"read"},
ResourceType: "did",
RequiresAuth: false,
Sensitive: false,
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "did:write",
Description: "Update DID documents",
UCANActions: []string{"read", "write"},
ResourceType: "did",
RequiresAuth: true,
Sensitive: true,
ParentScope: "did:read",
})
// DWN (Decentralized Web Node) scopes
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "dwn:read",
Description: "Read data from DWN",
UCANActions: []string{"read"},
ResourceType: "dwn",
RequiresAuth: true,
Sensitive: false,
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "dwn:write",
Description: "Write data to DWN",
UCANActions: []string{"read", "write"},
ResourceType: "dwn",
RequiresAuth: true,
Sensitive: true,
ParentScope: "dwn:read",
})
_ = m.RegisterScope(&OAuth2ScopeDefinition{
Name: "dwn:admin",
Description: "Full administrative access to DWN",
UCANActions: []string{"read", "write", "admin", "protocols-configure"},
ResourceType: "dwn",
RequiresAuth: true,
Sensitive: true,
ParentScope: "dwn:write",
ChildScopes: []string{"dwn:read", "dwn:write"},
})
}
// RegisterScope registers a new scope definition
func (m *ScopeMapper) RegisterScope(scope *OAuth2ScopeDefinition) error {
if scope.Name == "" {
return fmt.Errorf("scope name is required")
}
m.scopeDefinitions[scope.Name] = scope
// Create UCAN template for this scope
m.createUCANTemplate(scope)
return nil
}
// GetScope retrieves a scope definition
func (m *ScopeMapper) GetScope(name string) (*OAuth2ScopeDefinition, bool) {
scope, exists := m.scopeDefinitions[name]
return scope, exists
}
// MapToUCAN maps OAuth scopes to UCAN attenuations
func (m *ScopeMapper) MapToUCAN(
scopes []string,
userDID string,
clientID string,
resourceContext map[string]string,
) []ucan.Attenuation {
attenuations := []ucan.Attenuation{}
for _, scopeName := range scopes {
scope, exists := m.scopeDefinitions[scopeName]
if !exists {
continue
}
// Create attenuation for this scope
attenuation := m.createAttenuation(scope, userDID, clientID, resourceContext)
attenuations = append(attenuations, attenuation)
// Add child scope attenuations if this is a parent scope
for _, childScope := range scope.ChildScopes {
if childDef, exists := m.scopeDefinitions[childScope]; exists {
childAttenuation := m.createAttenuation(
childDef,
userDID,
clientID,
resourceContext,
)
attenuations = append(attenuations, childAttenuation)
}
}
}
return attenuations
}
// ValidateScopes validates that the requested scopes are valid
func (m *ScopeMapper) ValidateScopes(scopes []string) error {
for _, scope := range scopes {
if _, exists := m.scopeDefinitions[scope]; !exists {
return fmt.Errorf("invalid scope: %s", scope)
}
}
return nil
}
// GetScopeDescriptions returns human-readable descriptions for scopes
func (m *ScopeMapper) GetScopeDescriptions(scopes []string) map[string]string {
descriptions := make(map[string]string)
for _, scope := range scopes {
if def, exists := m.scopeDefinitions[scope]; exists {
descriptions[scope] = def.Description
}
}
return descriptions
}
// GetSensitiveScopes returns only the sensitive scopes from a list
func (m *ScopeMapper) GetSensitiveScopes(scopes []string) []string {
sensitive := []string{}
for _, scope := range scopes {
if def, exists := m.scopeDefinitions[scope]; exists && def.Sensitive {
sensitive = append(sensitive, scope)
}
}
return sensitive
}
// IsHierarchicalScope checks if one scope includes another
func (m *ScopeMapper) IsHierarchicalScope(parentScope, childScope string) bool {
parent, exists := m.scopeDefinitions[parentScope]
if !exists {
return false
}
// Check direct children
for _, child := range parent.ChildScopes {
if child == childScope {
return true
}
// Recursive check
if m.IsHierarchicalScope(child, childScope) {
return true
}
}
return false
}
// Private helper methods
func (m *ScopeMapper) createUCANTemplate(scope *OAuth2ScopeDefinition) {
// Create a template attenuation for this scope
capability := &ucan.SimpleCapability{
Action: strings.Join(scope.UCANActions, ","),
}
resource := &SimpleResource{
Scheme: scope.ResourceType,
Value: "{resource_id}",
}
m.ucanTemplates[scope.Name] = &ucan.Attenuation{
Capability: capability,
Resource: resource,
}
}
func (m *ScopeMapper) createAttenuation(
scope *OAuth2ScopeDefinition,
userDID string,
clientID string,
resourceContext map[string]string,
) ucan.Attenuation {
// Create capability based on scope actions
var capability ucan.Capability
if len(scope.UCANActions) == 1 {
capability = &ucan.SimpleCapability{
Action: scope.UCANActions[0],
}
} else {
capability = &ucan.MultiCapability{
Actions: scope.UCANActions,
}
}
// Create resource based on scope type and context
resourceValue := m.resolveResourceValue(scope, userDID, resourceContext)
resource := &SimpleResource{
Scheme: scope.ResourceType,
Value: resourceValue,
}
return ucan.Attenuation{
Capability: capability,
Resource: resource,
}
}
func (m *ScopeMapper) resolveResourceValue(
scope *OAuth2ScopeDefinition,
userDID string,
context map[string]string,
) string {
switch scope.ResourceType {
case "did":
return fmt.Sprintf("did:sonr:%s", userDID)
case "vault":
if vaultAddr, exists := context["vault_address"]; exists {
return vaultAddr
}
return fmt.Sprintf("vault:%s", userDID)
case "service":
if serviceID, exists := context["service_id"]; exists {
return serviceID
}
return "service:*"
case "dwn":
if dwnID, exists := context["dwn_id"]; exists {
return dwnID
}
return fmt.Sprintf("dwn:%s", userDID)
default:
return "*"
}
}
// SimpleResource implements the Resource interface for OAuth scopes
type SimpleResource struct {
Scheme string
Value string
}
func (r *SimpleResource) GetScheme() string {
return r.Scheme
}
func (r *SimpleResource) GetValue() string {
return r.Value
}
func (r *SimpleResource) GetURI() string {
return fmt.Sprintf("%s://%s", r.Scheme, r.Value)
}
func (r *SimpleResource) Matches(other ucan.Resource) bool {
if r.Scheme != other.GetScheme() {
return false
}
// Wildcard matching
if r.Value == "*" || other.GetValue() == "*" {
return true
}
// Exact match
return r.Value == other.GetValue()
}
+449
View File
@@ -0,0 +1,449 @@
package handlers
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"encoding/hex"
"fmt"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
)
const (
// PKCE challenge methods
PKCEMethodS256 = "S256"
PKCEMethodPlain = "plain"
)
// PKCEValidator handles PKCE validation
type PKCEValidator struct{}
// NewPKCEValidator creates a new PKCE validator
func NewPKCEValidator() *PKCEValidator {
return &PKCEValidator{}
}
// GeneratePKCEPair generates a PKCE verifier and challenge pair
func (p *PKCEValidator) GeneratePKCEPair() (verifier, challenge string, err error) {
// Generate cryptographically secure verifier (43-128 characters)
verifierBytes := make([]byte, 32)
if _, err := rand.Read(verifierBytes); err != nil {
return "", "", fmt.Errorf("failed to generate PKCE verifier: %w", err)
}
verifier = base64.RawURLEncoding.EncodeToString(verifierBytes)
// Create S256 challenge
challenge = p.CreateChallenge(verifier, PKCEMethodS256)
return verifier, challenge, nil
}
// CreateChallenge creates a PKCE challenge from a verifier
func (p *PKCEValidator) CreateChallenge(verifier, method string) string {
switch method {
case PKCEMethodS256:
h := sha256.Sum256([]byte(verifier))
return base64.RawURLEncoding.EncodeToString(h[:])
case PKCEMethodPlain:
return verifier
default:
// Default to S256 for security
h := sha256.Sum256([]byte(verifier))
return base64.RawURLEncoding.EncodeToString(h[:])
}
}
// Validate validates a PKCE verifier against a challenge
func (p *PKCEValidator) Validate(verifier, challenge, method string) bool {
if method == "" {
method = PKCEMethodPlain
}
computed := p.CreateChallenge(verifier, method)
return subtle.ConstantTimeCompare([]byte(computed), []byte(challenge)) == 1
}
// computePKCEChallenge computes a PKCE challenge (helper function)
func computePKCEChallenge(verifier, method string) string {
validator := NewPKCEValidator()
return validator.CreateChallenge(verifier, method)
}
// CSRFProtector handles CSRF protection
type CSRFProtector struct {
secret []byte
}
// NewCSRFProtector creates a new CSRF protector
func NewCSRFProtector(secret string) *CSRFProtector {
return &CSRFProtector{
secret: []byte(secret),
}
}
// GenerateToken generates a CSRF token
func (c *CSRFProtector) GenerateToken(sessionID string) (string, error) {
// Create a unique token tied to the session
h := hmac.New(sha256.New, c.secret)
h.Write([]byte(sessionID))
h.Write([]byte(time.Now().Format(time.RFC3339)))
tokenBytes := h.Sum(nil)
return base64.RawURLEncoding.EncodeToString(tokenBytes), nil
}
// ValidateToken validates a CSRF token
func (c *CSRFProtector) ValidateToken(token, sessionID string) bool {
// Decode the token
tokenBytes, err := base64.RawURLEncoding.DecodeString(token)
if err != nil {
return false
}
// Recreate the expected token
h := hmac.New(sha256.New, c.secret)
h.Write([]byte(sessionID))
// Note: In production, you'd want to include time validation
// and possibly store tokens with expiration
expectedBytes := h.Sum(nil)[:len(tokenBytes)]
return hmac.Equal(tokenBytes, expectedBytes)
}
// StateValidator validates OAuth state parameters
type StateValidator struct {
states map[string]*StateEntry
}
// StateEntry represents a stored state parameter
type StateEntry struct {
Value string
ClientID string
ExpiresAt time.Time
Used bool
}
// NewStateValidator creates a new state validator
func NewStateValidator() *StateValidator {
validator := &StateValidator{
states: make(map[string]*StateEntry),
}
// Start cleanup routine
go validator.cleanup()
return validator
}
// GenerateState generates a secure state parameter
func (s *StateValidator) GenerateState() (string, error) {
stateBytes := make([]byte, 32)
if _, err := rand.Read(stateBytes); err != nil {
return "", fmt.Errorf("failed to generate state: %w", err)
}
return base64.RawURLEncoding.EncodeToString(stateBytes), nil
}
// StoreState stores a state parameter for validation
func (s *StateValidator) StoreState(state, clientID string) {
s.states[state] = &StateEntry{
Value: state,
ClientID: clientID,
ExpiresAt: time.Now().Add(10 * time.Minute),
Used: false,
}
}
// ValidateState validates and consumes a state parameter
func (s *StateValidator) ValidateState(state, clientID string) bool {
entry, exists := s.states[state]
if !exists {
return false
}
// Check if expired
if time.Now().After(entry.ExpiresAt) {
delete(s.states, state)
return false
}
// Check if already used
if entry.Used {
return false
}
// Check client ID matches
if entry.ClientID != clientID {
return false
}
// Mark as used
entry.Used = true
return true
}
// cleanup removes expired states
func (s *StateValidator) cleanup() {
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for range ticker.C {
now := time.Now()
for state, entry := range s.states {
if now.After(entry.ExpiresAt) {
delete(s.states, state)
}
}
}
}
// JWTClientAuthenticator handles JWT client authentication
type JWTClientAuthenticator struct {
clientRegistry *ClientRegistry
}
// NewJWTClientAuthenticator creates a new JWT client authenticator
func NewJWTClientAuthenticator(registry *ClientRegistry) *JWTClientAuthenticator {
return &JWTClientAuthenticator{
clientRegistry: registry,
}
}
// ValidateClientAssertion validates a JWT client assertion
func (j *JWTClientAuthenticator) ValidateClientAssertion(
assertion, expectedAudience string,
) (*OAuth2Client, error) {
// Parse the JWT without verification first to get the claims
token, _, err := jwt.NewParser().ParseUnverified(assertion, jwt.MapClaims{})
if err != nil {
return nil, fmt.Errorf("failed to parse client assertion: %w", err)
}
claims, ok := token.Claims.(jwt.MapClaims)
if !ok {
return nil, fmt.Errorf("invalid claims format")
}
// Extract client ID from issuer and subject
clientID, ok := claims["iss"].(string)
if !ok {
return nil, fmt.Errorf("missing issuer claim")
}
subject, ok := claims["sub"].(string)
if !ok || subject != clientID {
return nil, fmt.Errorf("issuer and subject must match")
}
// Validate audience
audience, ok := claims["aud"].(string)
if !ok || audience != expectedAudience {
return nil, fmt.Errorf("invalid audience")
}
// Validate expiration
if exp, ok := claims["exp"].(float64); ok {
if time.Now().Unix() > int64(exp) {
return nil, fmt.Errorf("assertion expired")
}
} else {
return nil, fmt.Errorf("missing expiration")
}
// Validate not before
if nbf, ok := claims["nbf"].(float64); ok {
if time.Now().Unix() < int64(nbf) {
return nil, fmt.Errorf("assertion not yet valid")
}
}
// Validate issued at
if iat, ok := claims["iat"].(float64); ok {
// Check that the assertion is not too old (5 minutes max)
if time.Now().Unix()-int64(iat) > 300 {
return nil, fmt.Errorf("assertion too old")
}
}
// Validate JTI for replay protection
if jti, ok := claims["jti"].(string); !ok || jti == "" {
return nil, fmt.Errorf("missing jti claim")
}
// TODO: Store and check JTI to prevent replay attacks
// Get the client
client, err := j.clientRegistry.GetClient(clientID)
if err != nil {
return nil, fmt.Errorf("unknown client: %w", err)
}
// TODO: Verify the JWT signature using the client's registered public key
// This requires storing client public keys in the registry
return client, nil
}
// SecureTokenGenerator generates cryptographically secure tokens
type SecureTokenGenerator struct {
entropy int // bits of entropy
}
// NewSecureTokenGenerator creates a new secure token generator
func NewSecureTokenGenerator(entropyBits int) *SecureTokenGenerator {
if entropyBits < 128 {
entropyBits = 256 // Default to 256 bits for security
}
return &SecureTokenGenerator{
entropy: entropyBits,
}
}
// GenerateToken generates a secure random token
func (g *SecureTokenGenerator) GenerateToken() (string, error) {
bytes := make([]byte, g.entropy/8)
if _, err := rand.Read(bytes); err != nil {
return "", fmt.Errorf("failed to generate secure token: %w", err)
}
return base64.RawURLEncoding.EncodeToString(bytes), nil
}
// GenerateHexToken generates a secure random token in hex format
func (g *SecureTokenGenerator) GenerateHexToken() (string, error) {
bytes := make([]byte, g.entropy/8)
if _, err := rand.Read(bytes); err != nil {
return "", fmt.Errorf("failed to generate secure token: %w", err)
}
return hex.EncodeToString(bytes), nil
}
// RateLimiter implements rate limiting for OAuth endpoints
type RateLimiter struct {
attempts map[string][]time.Time
maxAttempts int
window time.Duration
}
// NewRateLimiter creates a new rate limiter
func NewRateLimiter(maxAttempts int, window time.Duration) *RateLimiter {
limiter := &RateLimiter{
attempts: make(map[string][]time.Time),
maxAttempts: maxAttempts,
window: window,
}
// Start cleanup routine
go limiter.cleanup()
return limiter
}
// Allow checks if a request should be allowed
func (r *RateLimiter) Allow(key string) bool {
now := time.Now()
windowStart := now.Add(-r.window)
// Get attempts for this key
attempts := r.attempts[key]
// Filter out attempts outside the window
validAttempts := []time.Time{}
for _, attempt := range attempts {
if attempt.After(windowStart) {
validAttempts = append(validAttempts, attempt)
}
}
// Check if under limit
if len(validAttempts) >= r.maxAttempts {
return false
}
// Add this attempt
validAttempts = append(validAttempts, now)
r.attempts[key] = validAttempts
return true
}
// cleanup removes old entries
func (r *RateLimiter) cleanup() {
ticker := time.NewTicker(5 * time.Minute)
defer ticker.Stop()
for range ticker.C {
now := time.Now()
windowStart := now.Add(-r.window)
for key, attempts := range r.attempts {
validAttempts := []time.Time{}
for _, attempt := range attempts {
if attempt.After(windowStart) {
validAttempts = append(validAttempts, attempt)
}
}
if len(validAttempts) == 0 {
delete(r.attempts, key)
} else {
r.attempts[key] = validAttempts
}
}
}
}
// OriginValidator validates request origins for CORS
type OriginValidator struct {
allowedOrigins map[string]bool
allowSubdomains bool
}
// NewOriginValidator creates a new origin validator
func NewOriginValidator(origins []string, allowSubdomains bool) *OriginValidator {
validator := &OriginValidator{
allowedOrigins: make(map[string]bool),
allowSubdomains: allowSubdomains,
}
for _, origin := range origins {
validator.allowedOrigins[origin] = true
}
return validator
}
// IsAllowed checks if an origin is allowed
func (o *OriginValidator) IsAllowed(origin string) bool {
// Direct match
if o.allowedOrigins[origin] {
return true
}
// Check subdomain matching if enabled
if o.allowSubdomains {
for allowed := range o.allowedOrigins {
if o.isSubdomainOf(origin, allowed) {
return true
}
}
}
return false
}
// isSubdomainOf checks if origin is a subdomain of allowed
func (o *OriginValidator) isSubdomainOf(origin, allowed string) bool {
// Simple subdomain check
// In production, use proper URL parsing
if strings.HasPrefix(allowed, "*.") {
domain := strings.TrimPrefix(allowed, "*")
return strings.HasSuffix(origin, domain)
}
return false
}
+650
View File
@@ -0,0 +1,650 @@
package handlers
import (
"context"
"encoding/json"
"fmt"
"net/http"
"strings"
"time"
"github.com/sonr-io/sonr/crypto/ucan"
)
// TokenExchangeHandler implements RFC 8693 OAuth 2.0 Token Exchange
type TokenExchangeHandler struct {
delegator *UCANDelegator
signer *BlockchainUCANSigner
tokenStore TokenStore
clientStore ClientStore
}
// TokenStore interface for token persistence
type TokenStore interface {
GetToken(ctx context.Context, tokenID string) (*StoredToken, error)
StoreToken(ctx context.Context, token *StoredToken) error
RevokeToken(ctx context.Context, tokenID string) error
}
// ClientStore interface for OAuth client information
type ClientStore interface {
GetClient(ctx context.Context, clientID string) (*OAuth2Client, error)
ValidateClientCredentials(ctx context.Context, clientID, clientSecret string) error
}
// StoredToken represents a stored OAuth token
type StoredToken struct {
TokenID string `json:"token_id"`
TokenType string `json:"token_type"`
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token,omitempty"`
ExpiresAt time.Time `json:"expires_at"`
Scopes []string `json:"scopes"`
ClientID string `json:"client_id"`
UserDID string `json:"user_did,omitempty"`
UCANToken string `json:"ucan_token"`
}
// TokenExchangeRequest represents an RFC 8693 token exchange request
type TokenExchangeRequest struct {
GrantType string `json:"grant_type"`
Resource string `json:"resource,omitempty"`
Audience string `json:"audience,omitempty"`
Scope string `json:"scope,omitempty"`
RequestedTokenType string `json:"requested_token_type,omitempty"`
SubjectToken string `json:"subject_token"`
SubjectTokenType string `json:"subject_token_type"`
ActorToken string `json:"actor_token,omitempty"`
ActorTokenType string `json:"actor_token_type,omitempty"`
}
// TokenExchangeResponse represents an RFC 8693 token exchange response
type TokenExchangeResponse struct {
AccessToken string `json:"access_token"`
IssuedTokenType string `json:"issued_token_type"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in,omitempty"`
Scope string `json:"scope,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
}
// Token type identifiers from RFC 8693
const (
TokenTypeAccessToken = "urn:ietf:params:oauth:token-type:access_token"
TokenTypeRefreshToken = "urn:ietf:params:oauth:token-type:refresh_token"
TokenTypeIDToken = "urn:ietf:params:oauth:token-type:id_token"
TokenTypeSAML1 = "urn:ietf:params:oauth:token-type:saml1"
TokenTypeSAML2 = "urn:ietf:params:oauth:token-type:saml2"
TokenTypeJWT = "urn:ietf:params:oauth:token-type:jwt"
TokenTypeUCAN = "urn:x-oath:params:oauth:token-type:ucan"
)
// NewTokenExchangeHandler creates a new token exchange handler
func NewTokenExchangeHandler(
delegator *UCANDelegator,
signer *BlockchainUCANSigner,
tokenStore TokenStore,
clientStore ClientStore,
) *TokenExchangeHandler {
return &TokenExchangeHandler{
delegator: delegator,
signer: signer,
tokenStore: tokenStore,
clientStore: clientStore,
}
}
// HandleTokenExchange handles RFC 8693 token exchange requests
func (h *TokenExchangeHandler) HandleTokenExchange(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
return
}
// Parse request
var req TokenExchangeRequest
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
h.sendError(w, "invalid_request", "Failed to parse request body")
return
}
// Validate grant type
if req.GrantType != "urn:ietf:params:oauth:grant-type:token-exchange" {
h.sendError(w, "unsupported_grant_type", "Only token-exchange grant type is supported")
return
}
// Validate required parameters
if req.SubjectToken == "" || req.SubjectTokenType == "" {
h.sendError(w, "invalid_request", "Missing required parameters")
return
}
// Authenticate client
clientID, clientSecret, ok := r.BasicAuth()
if !ok {
h.sendError(w, "invalid_client", "Client authentication required")
return
}
ctx := r.Context()
if err := h.clientStore.ValidateClientCredentials(ctx, clientID, clientSecret); err != nil {
h.sendError(w, "invalid_client", "Client authentication failed")
return
}
// Get client information
client, err := h.clientStore.GetClient(ctx, clientID)
if err != nil {
h.sendError(w, "invalid_client", "Client not found")
return
}
// Process token exchange based on token types
response, err := h.processTokenExchange(ctx, &req, client)
if err != nil {
h.sendError(w, "invalid_request", err.Error())
return
}
// Send response
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
json.NewEncoder(w).Encode(response)
}
// processTokenExchange processes the token exchange based on token types
func (h *TokenExchangeHandler) processTokenExchange(
ctx context.Context,
req *TokenExchangeRequest,
client *OAuth2Client,
) (*TokenExchangeResponse, error) {
// Determine requested token type (default to access token)
requestedType := req.RequestedTokenType
if requestedType == "" {
requestedType = TokenTypeAccessToken
}
// Handle different subject token types
switch req.SubjectTokenType {
case TokenTypeAccessToken:
return h.exchangeAccessToken(ctx, req, client, requestedType)
case TokenTypeRefreshToken:
return h.exchangeRefreshToken(ctx, req, client, requestedType)
case TokenTypeJWT:
return h.exchangeJWT(ctx, req, client, requestedType)
case TokenTypeUCAN:
return h.exchangeUCAN(ctx, req, client, requestedType)
default:
return nil, fmt.Errorf("unsupported subject token type: %s", req.SubjectTokenType)
}
}
// exchangeAccessToken exchanges an access token for a new token
func (h *TokenExchangeHandler) exchangeAccessToken(
ctx context.Context,
req *TokenExchangeRequest,
client *OAuth2Client,
requestedType string,
) (*TokenExchangeResponse, error) {
// Retrieve the subject token
storedToken, err := h.tokenStore.GetToken(ctx, req.SubjectToken)
if err != nil {
return nil, fmt.Errorf("invalid subject token")
}
// Validate token hasn't expired
if time.Now().After(storedToken.ExpiresAt) {
return nil, fmt.Errorf("subject token has expired")
}
// Parse requested scopes (default to original scopes)
requestedScopes := storedToken.Scopes
if req.Scope != "" {
requestedScopes = strings.Split(req.Scope, " ")
// Validate requested scopes are subset of original
if !h.isScopeSubset(requestedScopes, storedToken.Scopes) {
return nil, fmt.Errorf("requested scopes exceed original token scopes")
}
}
// Determine audience (default to requested audience or client ID)
audience := req.Audience
if audience == "" {
// Try to extract DID from client metadata
if clientDID, ok := client.Metadata["client_did"]; ok {
audience = clientDID
} else {
audience = client.ClientID
}
}
// Create new UCAN delegation based on requested type
switch requestedType {
case TokenTypeUCAN:
return h.createUCANResponse(
ctx,
storedToken.UserDID,
audience,
requestedScopes,
storedToken.UCANToken,
)
case TokenTypeAccessToken:
return h.createAccessTokenResponse(ctx, storedToken.UserDID, audience, requestedScopes)
default:
return nil, fmt.Errorf("unsupported requested token type: %s", requestedType)
}
}
// exchangeRefreshToken exchanges a refresh token for new tokens
func (h *TokenExchangeHandler) exchangeRefreshToken(
ctx context.Context,
req *TokenExchangeRequest,
client *OAuth2Client,
requestedType string,
) (*TokenExchangeResponse, error) {
// Retrieve the refresh token
storedToken, err := h.tokenStore.GetToken(ctx, req.SubjectToken)
if err != nil {
return nil, fmt.Errorf("invalid refresh token")
}
// Validate it's actually a refresh token
if storedToken.TokenType != "refresh_token" {
return nil, fmt.Errorf("token is not a refresh token")
}
// Create new tokens with same scopes
clientDID := client.ClientID
if did, ok := client.Metadata["client_did"]; ok {
clientDID = did
}
return h.createAccessTokenResponse(ctx, storedToken.UserDID, clientDID, storedToken.Scopes)
}
// exchangeJWT exchanges a JWT for a UCAN token
func (h *TokenExchangeHandler) exchangeJWT(
ctx context.Context,
req *TokenExchangeRequest,
client *OAuth2Client,
requestedType string,
) (*TokenExchangeResponse, error) {
// Verify the JWT
ucanToken, err := h.signer.VerifySignature(req.SubjectToken)
if err != nil {
return nil, fmt.Errorf("invalid JWT: %w", err)
}
// Extract scopes from JWT claims
scopes := h.extractScopesFromUCAN(ucanToken)
// Create response based on requested type
clientDID := client.ClientID
if did, ok := client.Metadata["client_did"]; ok {
clientDID = did
}
switch requestedType {
case TokenTypeUCAN:
return h.createUCANResponse(ctx, ucanToken.Issuer, clientDID, scopes, req.SubjectToken)
case TokenTypeAccessToken:
return h.createAccessTokenResponse(ctx, ucanToken.Issuer, clientDID, scopes)
default:
return nil, fmt.Errorf("unsupported requested token type: %s", requestedType)
}
}
// exchangeUCAN exchanges a UCAN token for another token type
func (h *TokenExchangeHandler) exchangeUCAN(
ctx context.Context,
req *TokenExchangeRequest,
client *OAuth2Client,
requestedType string,
) (*TokenExchangeResponse, error) {
// Verify the UCAN token
ucanToken, err := h.signer.VerifySignature(req.SubjectToken)
if err != nil {
return nil, fmt.Errorf("invalid UCAN token: %w", err)
}
// Validate delegation chain if actor token is provided
if req.ActorToken != "" {
actorToken, err := h.signer.VerifySignature(req.ActorToken)
if err != nil {
return nil, fmt.Errorf("invalid actor token: %w", err)
}
// Validate actor can act on behalf of subject
if actorToken.Audience != ucanToken.Issuer {
return nil, fmt.Errorf("actor token audience doesn't match subject issuer")
}
}
// Extract scopes from UCAN
scopes := h.extractScopesFromUCAN(ucanToken)
// Handle impersonation/delegation if actor token is present
issuer := ucanToken.Issuer
clientDID := client.ClientID
if did, ok := client.Metadata["client_did"]; ok {
clientDID = did
}
if req.ActorToken != "" {
// Actor is performing action on behalf of subject
issuer = clientDID // Actor becomes the new issuer
}
// Create response based on requested type
switch requestedType {
case TokenTypeAccessToken:
return h.createAccessTokenResponse(ctx, issuer, clientDID, scopes)
case TokenTypeUCAN:
// Create delegated UCAN with proof chain
proofs := []ucan.Proof{ucan.Proof(req.SubjectToken)}
if req.ActorToken != "" {
proofs = append(proofs, ucan.Proof(req.ActorToken))
}
return h.createDelegatedUCANResponse(ctx, issuer, clientDID, scopes, proofs)
default:
return nil, fmt.Errorf("unsupported requested token type: %s", requestedType)
}
}
// createUCANResponse creates a UCAN token response
func (h *TokenExchangeHandler) createUCANResponse(
ctx context.Context,
issuer, audience string,
scopes []string,
proof string,
) (*TokenExchangeResponse, error) {
// Create UCAN token with delegation
ucanToken, err := h.delegator.CreateDelegation(
issuer,
audience,
scopes,
time.Now().Add(time.Hour),
)
if err != nil {
return nil, fmt.Errorf("failed to create UCAN token: %w", err)
}
// Add proof if provided
if proof != "" {
ucanToken.Proofs = []ucan.Proof{ucan.Proof(proof)}
}
// Sign the token
signedToken, err := h.signer.Sign(ucanToken)
if err != nil {
return nil, fmt.Errorf("failed to sign UCAN token: %w", err)
}
return &TokenExchangeResponse{
AccessToken: signedToken,
IssuedTokenType: TokenTypeUCAN,
TokenType: "Bearer",
ExpiresIn: 3600,
Scope: strings.Join(scopes, " "),
}, nil
}
// createDelegatedUCANResponse creates a delegated UCAN token with proof chain
func (h *TokenExchangeHandler) createDelegatedUCANResponse(
ctx context.Context,
issuer, audience string,
scopes []string,
proofs []ucan.Proof,
) (*TokenExchangeResponse, error) {
// Build resource context
resourceContext := map[string]string{
"delegation_type": "token_exchange",
"issued_at": fmt.Sprintf("%d", time.Now().Unix()),
}
// Map scopes to attenuations
attenuations := h.delegator.scopeMapper.MapToUCAN(scopes, issuer, audience, resourceContext)
// Create UCAN token with proof chain
ucanToken := &ucan.Token{
Issuer: issuer,
Audience: audience,
ExpiresAt: time.Now().Add(time.Hour).Unix(),
NotBefore: time.Now().Unix(),
Attenuations: attenuations,
Proofs: proofs,
Facts: []ucan.Fact{
{
Data: h.createTokenExchangeFact(scopes),
},
},
}
// Sign the token
signedToken, err := h.signer.Sign(ucanToken)
if err != nil {
return nil, fmt.Errorf("failed to sign delegated UCAN token: %w", err)
}
return &TokenExchangeResponse{
AccessToken: signedToken,
IssuedTokenType: TokenTypeUCAN,
TokenType: "Bearer",
ExpiresIn: 3600,
Scope: strings.Join(scopes, " "),
}, nil
}
// createAccessTokenResponse creates a standard OAuth access token response
func (h *TokenExchangeHandler) createAccessTokenResponse(
ctx context.Context,
userDID, clientID string,
scopes []string,
) (*TokenExchangeResponse, error) {
// Create UCAN-backed access token
ucanToken, err := h.delegator.CreateDelegation(
userDID,
clientID,
scopes,
time.Now().Add(time.Hour),
)
if err != nil {
return nil, fmt.Errorf("failed to create access token: %w", err)
}
// Generate token ID
tokenID := h.generateTokenID()
// Store token
storedToken := &StoredToken{
TokenID: tokenID,
TokenType: "access_token",
AccessToken: tokenID,
ExpiresAt: time.Now().Add(time.Hour),
Scopes: scopes,
ClientID: clientID,
UserDID: userDID,
UCANToken: ucanToken.Raw,
}
if err := h.tokenStore.StoreToken(ctx, storedToken); err != nil {
return nil, fmt.Errorf("failed to store token: %w", err)
}
// Generate refresh token
refreshTokenID := h.generateTokenID()
refreshToken := &StoredToken{
TokenID: refreshTokenID,
TokenType: "refresh_token",
RefreshToken: refreshTokenID,
ExpiresAt: time.Now().Add(30 * 24 * time.Hour), // 30 days
Scopes: scopes,
ClientID: clientID,
UserDID: userDID,
}
if err := h.tokenStore.StoreToken(ctx, refreshToken); err != nil {
// Non-fatal, continue without refresh token
refreshTokenID = ""
}
response := &TokenExchangeResponse{
AccessToken: tokenID,
IssuedTokenType: TokenTypeAccessToken,
TokenType: "Bearer",
ExpiresIn: 3600,
Scope: strings.Join(scopes, " "),
}
if refreshTokenID != "" {
response.RefreshToken = refreshTokenID
}
return response, nil
}
// extractScopesFromUCAN extracts OAuth scopes from UCAN attenuations
func (h *TokenExchangeHandler) extractScopesFromUCAN(token *ucan.Token) []string {
scopeMap := make(map[string]bool)
for _, att := range token.Attenuations {
scheme := att.Resource.GetScheme()
actions := att.Capability.GetActions()
// Map UCAN capabilities back to OAuth scopes
for _, action := range actions {
scope := h.mapUCANToScope(scheme, action)
if scope != "" {
scopeMap[scope] = true
}
}
}
scopes := make([]string, 0, len(scopeMap))
for scope := range scopeMap {
scopes = append(scopes, scope)
}
return scopes
}
// mapUCANToScope maps UCAN capability to OAuth scope
func (h *TokenExchangeHandler) mapUCANToScope(scheme, action string) string {
// Reverse mapping from UCAN to OAuth scopes
switch scheme {
case "vault":
switch action {
case "read":
return "vault:read"
case "write":
return "vault:write"
case "sign":
return "vault:sign"
case "*", "admin":
return "vault:admin"
}
case "service", "svc":
switch action {
case "read":
return "service:read"
case "write":
return "service:write"
case "*", "admin":
return "service:manage"
}
case "did":
switch action {
case "read":
return "did:read"
case "write", "update":
return "did:write"
}
case "dwn":
switch action {
case "read":
return "dwn:read"
case "write":
return "dwn:write"
}
}
// Default mapping
return fmt.Sprintf("%s:%s", scheme, action)
}
// isScopeSubset checks if requested scopes are subset of allowed scopes
func (h *TokenExchangeHandler) isScopeSubset(requested, allowed []string) bool {
allowedMap := make(map[string]bool)
for _, scope := range allowed {
allowedMap[scope] = true
}
for _, scope := range requested {
if !allowedMap[scope] {
// Check if parent scope is allowed
if !h.isParentScopeAllowed(scope, allowed) {
return false
}
}
}
return true
}
// isParentScopeAllowed checks if a parent scope grants the requested scope
func (h *TokenExchangeHandler) isParentScopeAllowed(requested string, allowed []string) bool {
for _, scope := range allowed {
if h.delegator.scopeMapper.IsHierarchicalScope(scope, requested) {
return true
}
}
return false
}
// createTokenExchangeFact creates a fact for token exchange
func (h *TokenExchangeHandler) createTokenExchangeFact(scopes []string) json.RawMessage {
fact := map[string]any{
"type": "token_exchange",
"scopes": scopes,
"issued_at": time.Now().Unix(),
"grant_type": "urn:ietf:params:oauth:grant-type:token-exchange",
}
data, _ := json.Marshal(fact)
return json.RawMessage(data)
}
// generateTokenID generates a unique token identifier
func (h *TokenExchangeHandler) generateTokenID() string {
// In production, use a proper UUID or random generator
return fmt.Sprintf("tok_%d_%s", time.Now().UnixNano(), h.randomString(16))
}
// randomString generates a random string of specified length
func (h *TokenExchangeHandler) randomString(length int) string {
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
result := make([]byte, length)
for i := range result {
result[i] = charset[time.Now().UnixNano()%int64(len(charset))]
}
return string(result)
}
// sendError sends an OAuth error response
func (h *TokenExchangeHandler) sendError(
w http.ResponseWriter,
errorCode, errorDescription string,
) {
w.Header().Set("Content-Type", "application/json")
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
w.WriteHeader(http.StatusBadRequest)
response := map[string]string{
"error": errorCode,
"error_description": errorDescription,
}
json.NewEncoder(w).Encode(response)
}
+269
View File
@@ -0,0 +1,269 @@
package handlers
import (
"encoding/json"
"time"
"github.com/sonr-io/sonr/crypto/ucan"
)
// OAuth2Client represents a registered OAuth2 client application
type OAuth2Client struct {
ClientID string `json:"client_id"`
ClientSecret string `json:"client_secret,omitempty"`
ClientType string `json:"client_type"` // "confidential" or "public"
RedirectURIs []string `json:"redirect_uris"`
AllowedScopes []string `json:"allowed_scopes"`
AllowedGrants []string `json:"allowed_grants"`
TokenLifetime time.Duration `json:"token_lifetime"`
RequirePKCE bool `json:"require_pkce"`
TrustedClient bool `json:"trusted_client"`
RequiresConsent bool `json:"requires_consent"`
Metadata map[string]string `json:"metadata"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// OAuth2AuthorizationRequest represents an OAuth2 authorization request
type OAuth2AuthorizationRequest struct {
ResponseType string `json:"response_type" form:"response_type" query:"response_type"`
ClientID string `json:"client_id" form:"client_id" query:"client_id"`
RedirectURI string `json:"redirect_uri" form:"redirect_uri" query:"redirect_uri"`
Scope string `json:"scope" form:"scope" query:"scope"`
State string `json:"state" form:"state" query:"state"`
CodeChallenge string `json:"code_challenge" form:"code_challenge" query:"code_challenge"`
CodeChallengeMethod string `json:"code_challenge_method" form:"code_challenge_method" query:"code_challenge_method"`
Nonce string `json:"nonce" form:"nonce" query:"nonce"`
Prompt string `json:"prompt" form:"prompt" query:"prompt"`
MaxAge int `json:"max_age" form:"max_age" query:"max_age"`
LoginHint string `json:"login_hint" form:"login_hint" query:"login_hint"`
}
// OAuth2TokenRequest represents an OAuth2 token exchange request
type OAuth2TokenRequest struct {
GrantType string `json:"grant_type" form:"grant_type"`
Code string `json:"code" form:"code"`
RedirectURI string `json:"redirect_uri" form:"redirect_uri"`
ClientID string `json:"client_id" form:"client_id"`
ClientSecret string `json:"client_secret" form:"client_secret"`
CodeVerifier string `json:"code_verifier" form:"code_verifier"`
RefreshToken string `json:"refresh_token" form:"refresh_token"`
Scope string `json:"scope" form:"scope"`
Username string `json:"username" form:"username"`
Password string `json:"password" form:"password"`
Assertion string `json:"assertion" form:"assertion"`
ClientAssertion string `json:"client_assertion" form:"client_assertion"`
ClientAssertionType string `json:"client_assertion_type" form:"client_assertion_type"`
}
// OAuth2TokenResponse represents an OAuth2 token response
type OAuth2TokenResponse struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
RefreshToken string `json:"refresh_token,omitempty"`
Scope string `json:"scope,omitempty"`
IDToken string `json:"id_token,omitempty"`
UCANToken string `json:"ucan_token,omitempty"` // UCAN delegation token
}
// OAuth2ErrorResponse represents an OAuth2 error response
type OAuth2ErrorResponse struct {
Error string `json:"error"`
ErrorDescription string `json:"error_description,omitempty"`
ErrorURI string `json:"error_uri,omitempty"`
State string `json:"state,omitempty"`
}
// OAuth2AuthorizationCode represents an authorization code with UCAN context
type OAuth2AuthorizationCode struct {
Code string `json:"code"`
ClientID string `json:"client_id"`
UserDID string `json:"user_did"`
RedirectURI string `json:"redirect_uri"`
Scopes []string `json:"scopes"`
State string `json:"state"`
Nonce string `json:"nonce"`
CodeChallenge string `json:"code_challenge"`
CodeChallengeMethod string `json:"code_challenge_method"`
ExpiresAt time.Time `json:"expires_at"`
Used bool `json:"used"`
UCANContext *UCANAuthContext `json:"ucan_context"`
}
// UCANAuthContext holds UCAN-related data for authorization
type UCANAuthContext struct {
VaultAddress string `json:"vault_address"`
EnclaveDataCID string `json:"enclave_data_cid"`
DIDDocument json.RawMessage `json:"did_document"`
Capabilities []string `json:"capabilities"`
}
// OAuth2AccessToken represents an access token with embedded UCAN
type OAuth2AccessToken struct {
Token string `json:"token"`
UserDID string `json:"user_did"`
ClientID string `json:"client_id"`
Scopes []string `json:"scopes"`
ExpiresAt time.Time `json:"expires_at"`
IssuedAt time.Time `json:"issued_at"`
UCANToken *ucan.Token `json:"ucan_token"`
SessionID string `json:"session_id"`
TokenType string `json:"token_type"`
}
// OAuth2RefreshToken represents a refresh token
type OAuth2RefreshToken struct {
Token string `json:"token"`
AccessToken string `json:"access_token"`
ClientID string `json:"client_id"`
UserDID string `json:"user_did"`
Scopes []string `json:"scopes"`
ExpiresAt time.Time `json:"expires_at"`
IssuedAt time.Time `json:"issued_at"`
RotationCount int `json:"rotation_count"`
}
// OAuth2Session extends OIDCSession with OAuth2-specific fields
type OAuth2Session struct {
SessionID string `json:"session_id"`
UserDID string `json:"user_did"`
ClientID string `json:"client_id"`
Scopes []string `json:"scopes"`
AccessToken string `json:"access_token"`
RefreshToken string `json:"refresh_token"`
IDToken string `json:"id_token,omitempty"`
UCANToken *ucan.Token `json:"ucan_token"`
ConsentGiven bool `json:"consent_given"`
ConsentScopes []string `json:"consent_scopes"`
ExpiresAt time.Time `json:"expires_at"`
CreatedAt time.Time `json:"created_at"`
UpdatedAt time.Time `json:"updated_at"`
}
// OAuth2ConsentRequest represents a consent request for a client
type OAuth2ConsentRequest struct {
UserDID string `json:"user_did"`
ClientID string `json:"client_id"`
RequestedScopes []string `json:"requested_scopes"`
ConsentID string `json:"consent_id"`
CreatedAt time.Time `json:"created_at"`
}
// OAuth2ConsentResponse represents a user's consent response
type OAuth2ConsentResponse struct {
ConsentID string `json:"consent_id"`
Approved bool `json:"approved"`
ApprovedScopes []string `json:"approved_scopes"`
RememberChoice bool `json:"remember_choice"`
}
// OAuth2ScopeDefinition defines an OAuth scope and its UCAN mapping
type OAuth2ScopeDefinition struct {
Name string `json:"name"`
Description string `json:"description"`
UCANActions []string `json:"ucan_actions"`
ResourceType string `json:"resource_type"`
RequiresAuth bool `json:"requires_auth"`
Sensitive bool `json:"sensitive"`
ParentScope string `json:"parent_scope,omitempty"`
ChildScopes []string `json:"child_scopes,omitempty"`
Metadata map[string]any `json:"metadata"`
}
// OAuth2ClientRegistrationRequest represents a dynamic client registration request
type OAuth2ClientRegistrationRequest struct {
ClientName string `json:"client_name"`
ClientType string `json:"client_type"` // "confidential" or "public"
RedirectURIs []string `json:"redirect_uris"`
GrantTypes []string `json:"grant_types,omitempty"`
ResponseTypes []string `json:"response_types,omitempty"`
Scopes string `json:"scopes,omitempty"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
LogoURI string `json:"logo_uri,omitempty"`
ClientURI string `json:"client_uri,omitempty"`
PolicyURI string `json:"policy_uri,omitempty"`
TOSUri string `json:"tos_uri,omitempty"`
Contacts []string `json:"contacts,omitempty"`
}
// OAuth2ClientRegistrationResponse represents a successful client registration
type OAuth2ClientRegistrationResponse struct {
ClientID string `json:"client_id"`
ClientSecret string `json:"client_secret,omitempty"`
ClientIDIssuedAt int64 `json:"client_id_issued_at,omitempty"`
ClientSecretExpiresAt int64 `json:"client_secret_expires_at,omitempty"`
RegistrationAccessToken string `json:"registration_access_token,omitempty"`
RegistrationClientURI string `json:"registration_client_uri,omitempty"`
}
// OAuth2IntrospectionRequest represents a token introspection request
type OAuth2IntrospectionRequest struct {
Token string `json:"token" form:"token"`
TokenTypeHint string `json:"token_type_hint" form:"token_type_hint"`
ClientID string `json:"client_id" form:"client_id"`
ClientSecret string `json:"client_secret" form:"client_secret"`
}
// OAuth2IntrospectionResponse represents a token introspection response
type OAuth2IntrospectionResponse struct {
Active bool `json:"active"`
Scope string `json:"scope,omitempty"`
ClientID string `json:"client_id,omitempty"`
Username string `json:"username,omitempty"`
TokenType string `json:"token_type,omitempty"`
ExpiresAt int64 `json:"exp,omitempty"`
IssuedAt int64 `json:"iat,omitempty"`
NotBefore int64 `json:"nbf,omitempty"`
Subject string `json:"sub,omitempty"`
Audience []string `json:"aud,omitempty"`
Issuer string `json:"iss,omitempty"`
JWTID string `json:"jti,omitempty"`
UCANToken string `json:"ucan_token,omitempty"`
}
// OAuth2RevocationRequest represents a token revocation request
type OAuth2RevocationRequest struct {
Token string `json:"token" form:"token"`
TokenTypeHint string `json:"token_type_hint" form:"token_type_hint"`
ClientID string `json:"client_id" form:"client_id"`
ClientSecret string `json:"client_secret" form:"client_secret"`
}
// OAuth2Config extends OIDCConfig with OAuth2-specific endpoints
type OAuth2Config struct {
Issuer string `json:"issuer"`
AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"`
UserInfoEndpoint string `json:"userinfo_endpoint"`
JWKSEndpoint string `json:"jwks_uri"`
RegistrationEndpoint string `json:"registration_endpoint"`
IntrospectionEndpoint string `json:"introspection_endpoint"`
RevocationEndpoint string `json:"revocation_endpoint"`
ScopesSupported []string `json:"scopes_supported"`
ResponseTypesSupported []string `json:"response_types_supported"`
ResponseModesSupported []string `json:"response_modes_supported,omitempty"`
GrantTypesSupported []string `json:"grant_types_supported"`
SubjectTypesSupported []string `json:"subject_types_supported"`
IDTokenSigningAlgValuesSupported []string `json:"id_token_signing_alg_values_supported"`
TokenEndpointAuthMethodsSupported []string `json:"token_endpoint_auth_methods_supported"`
TokenEndpointAuthSigningAlgValuesSupported []string `json:"token_endpoint_auth_signing_alg_values_supported,omitempty"`
ClaimsSupported []string `json:"claims_supported"`
CodeChallengeMethodsSupported []string `json:"code_challenge_methods_supported"`
ServiceDocumentation string `json:"service_documentation,omitempty"`
UILocalesSupported []string `json:"ui_locales_supported,omitempty"`
OpPolicyURI string `json:"op_policy_uri,omitempty"`
OpTosURI string `json:"op_tos_uri,omitempty"`
UCANSupported bool `json:"ucan_supported"` // Custom field for UCAN support
}
// TokenValidationResult represents the result of token validation
type TokenValidationResult struct {
Valid bool `json:"valid"`
UserDID string `json:"user_did"`
ClientID string `json:"client_id"`
Scopes []string `json:"scopes"`
UCANToken *ucan.Token `json:"ucan_token,omitempty"`
ExpiresAt time.Time `json:"expires_at"`
Error string `json:"error,omitempty"`
}
+466
View File
@@ -0,0 +1,466 @@
package handlers
import (
"crypto/rand"
"crypto/sha256"
"encoding/base64"
"fmt"
"net/http"
"strings"
"sync"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/labstack/echo/v4"
)
// OIDCProvider manages OIDC operations
type OIDCProvider struct {
mu sync.RWMutex
codes map[string]*AuthorizationCode
sessions map[string]*OIDCSession
config any // TODO: Use bridge.OIDCProviderConfig
}
var (
oidcProvider = &OIDCProvider{
codes: make(map[string]*AuthorizationCode),
sessions: make(map[string]*OIDCSession),
}
codeExpiration = 10 * time.Minute
)
// SetOIDCConfig sets the OIDC provider configuration
func SetOIDCConfig(config any) {
oidcProvider.mu.Lock()
defer oidcProvider.mu.Unlock()
oidcProvider.config = config
}
// GetOIDCDiscovery returns OIDC discovery configuration
func GetOIDCDiscovery(c echo.Context) error {
config := &OIDCConfig{
Issuer: "https://localhost:8080",
AuthorizationEndpoint: "https://localhost:8080/oidc/authorize",
TokenEndpoint: "https://localhost:8080/oidc/token",
UserInfoEndpoint: "https://localhost:8080/oidc/userinfo",
JWKSEndpoint: "https://localhost:8080/oidc/jwks",
RevocationEndpoint: "https://localhost:8080/oidc/revoke",
IntrospectionEndpoint: "https://localhost:8080/oidc/introspect",
ScopesSupported: []string{
"openid", "profile", "email", "did", "vault", "offline_access",
},
ResponseTypesSupported: []string{
"code", "id_token", "code id_token",
},
GrantTypesSupported: []string{
"authorization_code", "refresh_token", "client_credentials",
},
SubjectTypesSupported: []string{
"public", "pairwise",
},
IDTokenSigningAlgValuesSupported: []string{
"ES256", "RS256",
},
TokenEndpointAuthMethodsSupported: []string{
"client_secret_post", "client_secret_basic", "none",
},
ClaimsSupported: []string{
"sub", "name", "preferred_username", "email", "email_verified",
"did", "vault_id", "updated_at",
},
CodeChallengeMethodsSupported: []string{
"S256", "plain",
},
}
return c.JSON(http.StatusOK, config)
}
// HandleOIDCAuthorization handles OIDC authorization requests
func HandleOIDCAuthorization(c echo.Context) error {
// Parse request parameters from query string for GET or form for POST
req := OIDCAuthorizationRequest{
ResponseType: c.QueryParam("response_type"),
ClientID: c.QueryParam("client_id"),
RedirectURI: c.QueryParam("redirect_uri"),
Scope: c.QueryParam("scope"),
State: c.QueryParam("state"),
Nonce: c.QueryParam("nonce"),
CodeChallenge: c.QueryParam("code_challenge"),
CodeChallengeMethod: c.QueryParam("code_challenge_method"),
}
// If no query params, try to bind from body (for POST requests)
if req.ResponseType == "" {
_ = c.Bind(&req) // Ignore bind errors and continue
}
// Validate required parameters
if req.ResponseType == "" || req.ClientID == "" || req.RedirectURI == "" || req.Scope == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Missing required parameters",
})
}
// Validate response type
if req.ResponseType != "code" && req.ResponseType != "id_token" &&
req.ResponseType != "code id_token" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "unsupported_response_type",
"error_description": "Response type not supported",
})
}
// Validate scope includes openid
if !strings.Contains(req.Scope, "openid") {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_scope",
"error_description": "Scope must include 'openid'",
})
}
// TODO: Validate client_id and redirect_uri against registered clients
// For now, assume user is authenticated (in production, redirect to login)
userDID := c.Get("user_did")
if userDID == nil {
// Redirect to WebAuthn authentication
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "authentication_required",
"error_description": "User authentication required",
})
}
// Generate authorization code
code := generateAuthorizationCode()
// Store authorization code
authCode := &AuthorizationCode{
Code: code,
ClientID: req.ClientID,
RedirectURI: req.RedirectURI,
UserDID: userDID.(string),
Scope: req.Scope,
Nonce: req.Nonce,
CodeChallenge: req.CodeChallenge,
CodeChallengeMethod: req.CodeChallengeMethod,
ExpiresAt: time.Now().Add(codeExpiration),
Used: false,
}
oidcProvider.mu.Lock()
oidcProvider.codes[code] = authCode
oidcProvider.mu.Unlock()
// Build redirect URL
redirectURL := fmt.Sprintf("%s?code=%s&state=%s", req.RedirectURI, code, req.State)
return c.Redirect(http.StatusFound, redirectURL)
}
// HandleOIDCToken handles OIDC token requests
func HandleOIDCToken(c echo.Context) error {
var req OIDCTokenRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Invalid token request",
})
}
// Handle different grant types
switch req.GrantType {
case "authorization_code":
return handleAuthorizationCodeGrant(c, &req)
case "refresh_token":
return handleRefreshTokenGrant(c, &req)
case "client_credentials":
return handleClientCredentialsGrant(c, &req)
default:
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "unsupported_grant_type",
"error_description": "Grant type not supported",
})
}
}
// handleAuthorizationCodeGrant processes authorization code grant
func handleAuthorizationCodeGrant(c echo.Context, req *OIDCTokenRequest) error {
if req.Code == "" || req.RedirectURI == "" || req.ClientID == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Missing required parameters",
})
}
// Retrieve and validate authorization code
oidcProvider.mu.Lock()
authCode, exists := oidcProvider.codes[req.Code]
if exists {
delete(oidcProvider.codes, req.Code) // Single use
}
oidcProvider.mu.Unlock()
if !exists {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_grant",
"error_description": "Invalid authorization code",
})
}
// Validate code hasn't expired
if time.Now().After(authCode.ExpiresAt) {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_grant",
"error_description": "Authorization code expired",
})
}
// Validate code hasn't been used
if authCode.Used {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_grant",
"error_description": "Authorization code already used",
})
}
// Validate client and redirect URI
if authCode.ClientID != req.ClientID || authCode.RedirectURI != req.RedirectURI {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_grant",
"error_description": "Invalid client or redirect URI",
})
}
// Validate PKCE if present
if authCode.CodeChallenge != "" {
if req.CodeVerifier == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_grant",
"error_description": "Code verifier required",
})
}
if !verifyPKCE(req.CodeVerifier, authCode.CodeChallenge, authCode.CodeChallengeMethod) {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_grant",
"error_description": "Invalid code verifier",
})
}
}
// Mark code as used
authCode.Used = true
// Generate tokens
accessToken := generateAccessToken(authCode.UserDID)
refreshToken := generateRefreshToken()
idToken, err := generateIDToken(authCode.UserDID, authCode.ClientID, authCode.Nonce)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "server_error",
"error_description": "Failed to generate ID token",
})
}
// Create session
session := &OIDCSession{
SessionID: generateSessionID(),
UserDID: authCode.UserDID,
ClientID: authCode.ClientID,
Scope: authCode.Scope,
Nonce: authCode.Nonce,
AccessToken: accessToken,
RefreshToken: refreshToken,
IDToken: idToken,
ExpiresAt: time.Now().Add(time.Hour),
CreatedAt: time.Now(),
}
oidcProvider.mu.Lock()
oidcProvider.sessions[accessToken] = session
oidcProvider.mu.Unlock()
// Return tokens
response := &OIDCTokenResponse{
AccessToken: accessToken,
TokenType: "Bearer",
ExpiresIn: 3600,
RefreshToken: refreshToken,
IDToken: idToken,
Scope: authCode.Scope,
}
return c.JSON(http.StatusOK, response)
}
// handleRefreshTokenGrant processes refresh token grant
func handleRefreshTokenGrant(c echo.Context, req *OIDCTokenRequest) error {
// TODO: Implement refresh token grant
return c.JSON(http.StatusNotImplemented, map[string]string{
"error": "unsupported_grant_type",
"error_description": "Refresh token grant not yet implemented",
})
}
// handleClientCredentialsGrant processes client credentials grant
func handleClientCredentialsGrant(c echo.Context, req *OIDCTokenRequest) error {
// TODO: Implement client credentials grant
return c.JSON(http.StatusNotImplemented, map[string]string{
"error": "unsupported_grant_type",
"error_description": "Client credentials grant not yet implemented",
})
}
// HandleOIDCUserInfo handles OIDC userinfo requests
func HandleOIDCUserInfo(c echo.Context) error {
// Extract access token from Authorization header
authHeader := c.Request().Header.Get("Authorization")
if authHeader == "" || !strings.HasPrefix(authHeader, "Bearer ") {
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "invalid_token",
"error_description": "Invalid access token",
})
}
accessToken := strings.TrimPrefix(authHeader, "Bearer ")
// Validate access token
oidcProvider.mu.RLock()
session, exists := oidcProvider.sessions[accessToken]
oidcProvider.mu.RUnlock()
if !exists {
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "invalid_token",
"error_description": "Invalid or expired access token",
})
}
// Check if token is expired
if time.Now().After(session.ExpiresAt) {
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "invalid_token",
"error_description": "Access token expired",
})
}
// TODO: Fetch actual user info from DID document or database
userInfo := &OIDCUserInfo{
Subject: session.UserDID,
PreferredUsername: "user", // TODO: Get from DID document
DID: session.UserDID,
UpdatedAt: time.Now().Unix(),
}
// Add additional claims based on scope
if strings.Contains(session.Scope, "profile") {
userInfo.Name = "User Name" // TODO: Get from DID document
}
if strings.Contains(session.Scope, "email") {
userInfo.Email = "user@example.com" // TODO: Get from DID document
userInfo.EmailVerified = true
}
if strings.Contains(session.Scope, "vault") {
userInfo.VaultID = "vault_" + session.UserDID // TODO: Get actual vault ID
}
return c.JSON(http.StatusOK, userInfo)
}
// HandleOIDCJWKS handles JWKS endpoint
func HandleOIDCJWKS(c echo.Context) error {
// TODO: Return actual public keys used for signing
jwks := &JWKSet{
Keys: []JWK{
{
KeyType: "EC",
Use: "sig",
KeyID: "1",
Algorithm: "ES256",
Curve: "P-256",
// TODO: Add actual public key coordinates
},
},
}
return c.JSON(http.StatusOK, jwks)
}
// Helper functions
func generateAuthorizationCode() string {
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
func generateAccessToken(userDID string) string {
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
func generateRefreshToken() string {
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
func generateSessionID() string {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
func generateIDToken(userDID, clientID, nonce string) (string, error) {
claims := &DIDAuthClaims{
RegisteredClaims: jwt.RegisteredClaims{
Subject: userDID,
Issuer: "https://localhost:8080",
Audience: []string{clientID},
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
},
DID: userDID,
}
if nonce != "" {
claims.Extra = map[string]any{
"nonce": nonce,
}
}
// TODO: Sign with actual signing key
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString([]byte("temporary-secret-key"))
}
func verifyPKCE(verifier, challenge, method string) bool {
var computed string
switch method {
case "S256":
h := sha256.Sum256([]byte(verifier))
computed = base64.RawURLEncoding.EncodeToString(h[:])
case "plain":
computed = verifier
default:
return false
}
return computed == challenge
}
+360
View File
@@ -0,0 +1,360 @@
package handlers
import (
"encoding/json"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/labstack/echo/v4"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
// TestOIDCDiscovery tests the OIDC discovery endpoint
func TestOIDCDiscovery(t *testing.T) {
e := echo.New()
req := httptest.NewRequest(http.MethodGet, "/.well-known/openid-configuration", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := GetOIDCDiscovery(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var config OIDCConfig
err = json.Unmarshal(rec.Body.Bytes(), &config)
assert.NoError(t, err)
// Verify required fields
assert.NotEmpty(t, config.Issuer)
assert.NotEmpty(t, config.AuthorizationEndpoint)
assert.NotEmpty(t, config.TokenEndpoint)
assert.NotEmpty(t, config.UserInfoEndpoint)
assert.NotEmpty(t, config.JWKSEndpoint)
assert.Contains(t, config.ScopesSupported, "openid")
assert.Contains(t, config.ResponseTypesSupported, "code")
assert.Contains(t, config.GrantTypesSupported, "authorization_code")
}
// TestOIDCAuthorizationFlow tests the authorization code flow
func TestOIDCAuthorizationFlow(t *testing.T) {
e := echo.New()
t.Run("ValidAuthorizationRequest", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "code")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid profile")
q.Set("state", "test-state")
q.Set("nonce", "test-nonce")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
// Set authenticated user context
c.Set("user_did", "did:sonr:testuser")
c.Set("authenticated", true)
err := HandleOIDCAuthorization(c)
assert.NoError(t, err)
// Should redirect with authorization code
assert.Equal(t, http.StatusFound, rec.Code)
location := rec.Header().Get("Location")
assert.Contains(t, location, "code=")
assert.Contains(t, location, "state=test-state")
})
t.Run("MissingRequiredParameters", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/authorize", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleOIDCAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
var errorResp map[string]string
err = json.Unmarshal(rec.Body.Bytes(), &errorResp)
assert.NoError(t, err)
assert.Equal(t, "invalid_request", errorResp["error"])
})
t.Run("InvalidResponseType", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "invalid")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleOIDCAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
})
}
// TestOIDCTokenExchange tests the token endpoint
func TestOIDCTokenExchange(t *testing.T) {
e := echo.New()
// Setup: Create an authorization code
code := "test-auth-code"
authCode := &AuthorizationCode{
Code: code,
ClientID: "test-client",
RedirectURI: "http://localhost:3000/callback",
UserDID: "did:sonr:testuser",
Scope: "openid profile",
ExpiresAt: time.Now().Add(10 * time.Minute),
CodeChallenge: "test-challenge",
CodeChallengeMethod: "S256",
}
oidcProvider.mu.Lock()
oidcProvider.codes[code] = authCode
oidcProvider.mu.Unlock()
t.Run("ValidTokenExchange", func(t *testing.T) {
body := strings.NewReader("grant_type=authorization_code&code=" + code +
"&redirect_uri=http://localhost:3000/callback&client_id=test-client" +
"&code_verifier=test-verifier")
req := httptest.NewRequest(http.MethodPost, "/oidc/token", body)
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
tokenReq := &OIDCTokenRequest{
GrantType: "authorization_code",
Code: code,
RedirectURI: "http://localhost:3000/callback",
ClientID: "test-client",
CodeVerifier: "test-verifier",
}
err := handleAuthorizationCodeGrant(c, tokenReq)
assert.NoError(t, err)
if rec.Code == http.StatusOK {
var tokenResp OIDCTokenResponse
err = json.Unmarshal(rec.Body.Bytes(), &tokenResp)
assert.NoError(t, err)
assert.NotEmpty(t, tokenResp.AccessToken)
assert.NotEmpty(t, tokenResp.IDToken)
assert.Equal(t, "Bearer", tokenResp.TokenType)
}
})
t.Run("ExpiredAuthorizationCode", func(t *testing.T) {
expiredCode := "expired-code"
expiredAuthCode := &AuthorizationCode{
Code: expiredCode,
ClientID: "test-client",
RedirectURI: "http://localhost:3000/callback",
UserDID: "did:sonr:testuser",
ExpiresAt: time.Now().Add(-1 * time.Hour), // Expired
}
oidcProvider.mu.Lock()
oidcProvider.codes[expiredCode] = expiredAuthCode
oidcProvider.mu.Unlock()
tokenReq := &OIDCTokenRequest{
GrantType: "authorization_code",
Code: expiredCode,
RedirectURI: "http://localhost:3000/callback",
ClientID: "test-client",
}
req := httptest.NewRequest(http.MethodPost, "/oidc/token", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := handleAuthorizationCodeGrant(c, tokenReq)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
})
}
// TestOIDCUserInfo tests the userinfo endpoint
func TestOIDCUserInfo(t *testing.T) {
e := echo.New()
// Setup: Create a session
accessToken := "test-access-token"
session := &OIDCSession{
SessionID: "test-session",
UserDID: "did:sonr:testuser",
ClientID: "test-client",
Scope: "openid profile email",
AccessToken: accessToken,
RefreshToken: "test-refresh-token",
ExpiresAt: time.Now().Add(1 * time.Hour),
CreatedAt: time.Now(),
}
oidcProvider.mu.Lock()
oidcProvider.sessions[accessToken] = session
oidcProvider.mu.Unlock()
t.Run("ValidUserInfoRequest", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/userinfo", nil)
req.Header.Set("Authorization", "Bearer "+accessToken)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleOIDCUserInfo(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var userInfo map[string]any
err = json.Unmarshal(rec.Body.Bytes(), &userInfo)
assert.NoError(t, err)
assert.Equal(t, "did:sonr:testuser", userInfo["sub"])
})
t.Run("InvalidAccessToken", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/userinfo", nil)
req.Header.Set("Authorization", "Bearer invalid-token")
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleOIDCUserInfo(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusUnauthorized, rec.Code)
})
t.Run("MissingAuthorizationHeader", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/userinfo", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleOIDCUserInfo(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusUnauthorized, rec.Code)
})
}
// TestPKCEFlow tests PKCE (Proof Key for Code Exchange) implementation
func TestPKCEFlow(t *testing.T) {
e := echo.New()
// Generate PKCE parameters
codeVerifier := "test-code-verifier-string-that-is-long-enough"
codeChallenge := "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM" // SHA256 of verifier
t.Run("AuthorizationWithPKCE", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/oidc/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "code")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid")
q.Set("code_challenge", codeChallenge)
q.Set("code_challenge_method", "S256")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
c.Set("user_did", "did:sonr:testuser")
c.Set("authenticated", true)
err := HandleOIDCAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusFound, rec.Code)
})
t.Run("TokenExchangeWithPKCE", func(t *testing.T) {
// Create auth code with PKCE
code := "pkce-auth-code"
authCode := &AuthorizationCode{
Code: code,
ClientID: "test-client",
RedirectURI: "http://localhost:3000/callback",
UserDID: "did:sonr:testuser",
CodeChallenge: codeChallenge,
CodeChallengeMethod: "S256",
ExpiresAt: time.Now().Add(10 * time.Minute),
}
oidcProvider.mu.Lock()
oidcProvider.codes[code] = authCode
oidcProvider.mu.Unlock()
tokenReq := &OIDCTokenRequest{
GrantType: "authorization_code",
Code: code,
RedirectURI: "http://localhost:3000/callback",
ClientID: "test-client",
CodeVerifier: codeVerifier,
}
req := httptest.NewRequest(http.MethodPost, "/oidc/token", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := handleAuthorizationCodeGrant(c, tokenReq)
require.NoError(t, err)
// With correct verifier, should succeed
if rec.Code != http.StatusOK {
t.Logf("Response: %s", rec.Body.String())
}
})
}
// TestRefreshTokenFlow tests refresh token functionality
func TestRefreshTokenFlow(t *testing.T) {
e := echo.New()
// Create initial session with refresh token
refreshToken := "test-refresh-token"
session := &OIDCSession{
SessionID: "test-session",
UserDID: "did:sonr:testuser",
ClientID: "test-client",
Scope: "openid profile offline_access",
AccessToken: "old-access-token",
RefreshToken: refreshToken,
ExpiresAt: time.Now().Add(-1 * time.Hour), // Expired access token
CreatedAt: time.Now().Add(-2 * time.Hour),
}
oidcProvider.mu.Lock()
oidcProvider.sessions[refreshToken] = session
oidcProvider.mu.Unlock()
t.Run("ValidRefreshToken", func(t *testing.T) {
tokenReq := &OIDCTokenRequest{
GrantType: "refresh_token",
RefreshToken: refreshToken,
ClientID: "test-client",
}
req := httptest.NewRequest(http.MethodPost, "/oidc/token", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := handleRefreshTokenGrant(c, tokenReq)
assert.NoError(t, err)
if rec.Code == http.StatusOK {
var tokenResp OIDCTokenResponse
err = json.Unmarshal(rec.Body.Bytes(), &tokenResp)
assert.NoError(t, err)
assert.NotEmpty(t, tokenResp.AccessToken)
assert.NotEqual(t, "old-access-token", tokenResp.AccessToken)
}
})
}
+384
View File
@@ -0,0 +1,384 @@
package handlers
import (
"crypto/rand"
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strings"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/labstack/echo/v4"
)
// SIOPProvider manages Self-Issued OpenID Provider operations
type SIOPProvider struct {
// No persistent state needed for SIOP as it's self-issued
}
var siopProvider = &SIOPProvider{}
// HandleSIOPAuthorization handles SIOP authorization requests
func HandleSIOPAuthorization(c echo.Context) error {
var req SIOPRequest
// Check if request was already parsed by HandleSIOPRequest
if parsedReq := c.Get("siop_request"); parsedReq != nil {
req = *(parsedReq.(*SIOPRequest))
} else {
// Parse request parameters from query string for GET or form for POST
req = SIOPRequest{
ResponseType: c.QueryParam("response_type"),
ClientID: c.QueryParam("client_id"),
RedirectURI: c.QueryParam("redirect_uri"),
Scope: c.QueryParam("scope"),
Nonce: c.QueryParam("nonce"),
State: c.QueryParam("state"),
}
// Parse claims if provided
if claimsStr := c.QueryParam("claims"); claimsStr != "" {
_ = json.Unmarshal([]byte(claimsStr), &req.Claims) // Ignore claims parsing errors
}
// If no query params, try to bind from body (for POST requests)
if req.ResponseType == "" {
_ = c.Bind(&req) // Ignore bind errors and continue with empty request
}
}
// Validate required parameters
if req.ResponseType == "" || req.ClientID == "" || req.RedirectURI == "" ||
req.Scope == "" || req.Nonce == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Missing required parameters",
})
}
// Validate response type (SIOP v2 supports id_token)
if req.ResponseType != "id_token" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "unsupported_response_type",
"error_description": "SIOP only supports id_token response type",
})
}
// Validate scope includes openid
if !strings.Contains(req.Scope, "openid") {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_scope",
"error_description": "Scope must include 'openid'",
})
}
// Get user DID from context (assumes user is authenticated via WebAuthn)
userDID := c.Get("user_did")
if userDID == nil {
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "authentication_required",
"error_description": "User authentication required for SIOP",
})
}
// Generate Self-Issued ID Token
idToken, err := generateSelfIssuedIDToken(
userDID.(string),
req.ClientID,
req.Nonce,
req.RedirectURI,
req.Claims,
)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "server_error",
"error_description": "Failed to generate self-issued ID token",
})
}
// Build response
response := &SIOPResponse{
IDToken: idToken,
State: req.State,
}
// If VP token is requested, include it
if req.Claims.VPToken != nil {
vpToken, err := generateVPToken(userDID.(string), req.Claims.VPToken)
if err == nil {
response.VPToken = vpToken
}
}
// Redirect back to client with response
redirectURL, err := buildSIOPRedirectURL(req.RedirectURI, response)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "server_error",
"error_description": "Failed to build redirect URL",
})
}
return c.Redirect(http.StatusFound, redirectURL)
}
// HandleSIOPRegistration handles dynamic client registration for SIOP
func HandleSIOPRegistration(c echo.Context) error {
var registration SIOPRegistration
if err := c.Bind(&registration); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Invalid registration request",
})
}
// Validate registration parameters
if registration.ClientName == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Client name is required",
})
}
// Generate client ID for the registration
clientID := generateClientID()
// Store registration (in production, persist this)
// For now, return the registration confirmation
response := map[string]any{
"client_id": clientID,
"client_name": registration.ClientName,
"client_purpose": registration.ClientPurpose,
"logo_uri": registration.LogoURI,
"subject_syntax_types_supported": []string{"did"},
"vp_formats": map[string]any{
"jwt_vp": map[string]any{
"alg": []string{"ES256", "EdDSA"},
},
"ldp_vp": map[string]any{
"proof_type": []string{"Ed25519Signature2018"},
},
},
}
return c.JSON(http.StatusOK, response)
}
// HandleSIOPMetadata returns SIOP provider metadata
func HandleSIOPMetadata(c echo.Context) error {
metadata := map[string]any{
"issuer": "https://self-issued.me/v2",
"authorization_endpoint": "openid://",
"response_types_supported": []string{"id_token"},
"scopes_supported": []string{
"openid",
"profile",
"email",
"did",
},
"subject_types_supported": []string{"pairwise"},
"id_token_signing_alg_values_supported": []string{"ES256", "EdDSA"},
"request_object_signing_alg_values_supported": []string{"ES256", "EdDSA"},
"subject_syntax_types_supported": []string{
"did:key",
"did:web",
"did:ion",
},
"vp_formats_supported": map[string]any{
"jwt_vp": map[string]any{
"alg_values_supported": []string{"ES256", "EdDSA"},
},
"ldp_vp": map[string]any{
"proof_type_values_supported": []string{"Ed25519Signature2018"},
},
},
}
return c.JSON(http.StatusOK, metadata)
}
// generateSelfIssuedIDToken generates a self-issued ID token
func generateSelfIssuedIDToken(
userDID, clientID, nonce, redirectURI string,
requestedClaims SIOPClaims,
) (string, error) {
// Create base claims
claims := jwt.MapClaims{
"iss": "https://self-issued.me/v2",
"sub": userDID,
"aud": clientID,
"exp": time.Now().Add(time.Hour).Unix(),
"iat": time.Now().Unix(),
"nonce": nonce,
"_sd_alg": "sha-256", // Selective disclosure algorithm
"sub_jwk": map[string]any{
// TODO: Include actual public key from DID document
"kty": "EC",
"crv": "P-256",
"x": "placeholder_x",
"y": "placeholder_y",
},
}
// Add requested claims from ID token
if requestedClaims.IDToken != nil {
for claimName := range requestedClaims.IDToken {
// Add claim based on request
switch claimName {
case "email":
claims["email"] = "user@example.com" // TODO: Get from DID document
claims["email_verified"] = true
case "name":
claims["name"] = "User Name" // TODO: Get from DID document
case "did":
claims["did"] = userDID
}
}
}
// Sign with user's DID key (simplified for now)
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString([]byte("temporary-siop-key"))
}
// generateVPToken generates a Verifiable Presentation token
func generateVPToken(userDID string, vpClaims map[string]ClaimRequest) (string, error) {
// Create VP structure
vp := map[string]any{
"@context": []string{
"https://www.w3.org/2018/credentials/v1",
},
"type": []string{"VerifiablePresentation"},
"holder": userDID,
"verifiableCredential": []any{
// TODO: Include actual verifiable credentials
},
}
// Convert to JWT
vpJSON, err := json.Marshal(vp)
if err != nil {
return "", err
}
// Sign VP (simplified for now)
claims := jwt.MapClaims{
"vp": string(vpJSON),
"iss": userDID,
"aud": "verifier",
"exp": time.Now().Add(time.Hour).Unix(),
"iat": time.Now().Unix(),
"nonce": generateNonce(),
}
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
return token.SignedString([]byte("temporary-vp-key"))
}
// buildSIOPRedirectURL builds the redirect URL with SIOP response
func buildSIOPRedirectURL(redirectURI string, response *SIOPResponse) (string, error) {
u, err := url.Parse(redirectURI)
if err != nil {
return "", err
}
// Add response parameters
q := u.Query()
q.Set("id_token", response.IDToken)
if response.VPToken != "" {
q.Set("vp_token", response.VPToken)
}
if response.State != "" {
q.Set("state", response.State)
}
u.RawQuery = q.Encode()
return u.String(), nil
}
// generateClientID generates a unique client ID
func generateClientID() string {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return fmt.Sprintf("siop_client_%s", base64.RawURLEncoding.EncodeToString(bytes))
}
// generateNonce generates a nonce for tokens
func generateNonce() string {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return ""
}
return base64.RawURLEncoding.EncodeToString(bytes)
}
// ValidateSIOPRequest validates an incoming SIOP request
func ValidateSIOPRequest(requestURI string) (*SIOPRequest, error) {
// Parse the request URI
u, err := url.Parse(requestURI)
if err != nil {
return nil, fmt.Errorf("invalid request URI: %w", err)
}
// Validate that it's a proper URI with scheme and host
if u.Scheme == "" && u.Host == "" {
return nil, fmt.Errorf("invalid request URI: missing scheme or host")
}
// Extract parameters
q := u.Query()
req := &SIOPRequest{
ResponseType: q.Get("response_type"),
ClientID: q.Get("client_id"),
RedirectURI: q.Get("redirect_uri"),
Scope: q.Get("scope"),
Nonce: q.Get("nonce"),
State: q.Get("state"),
}
// Parse claims if present
if claimsStr := q.Get("claims"); claimsStr != "" {
if err := json.Unmarshal([]byte(claimsStr), &req.Claims); err != nil {
return nil, fmt.Errorf("invalid claims parameter: %w", err)
}
}
// Parse registration if present
if regStr := q.Get("registration"); regStr != "" {
if err := json.Unmarshal([]byte(regStr), &req.Registration); err != nil {
return nil, fmt.Errorf("invalid registration parameter: %w", err)
}
}
return req, nil
}
// HandleSIOPRequest processes a complete SIOP request flow
func HandleSIOPRequest(c echo.Context) error {
// Get request URI from query parameter
requestURI := c.QueryParam("request_uri")
if requestURI == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": "Missing request_uri parameter",
})
}
// Validate and parse the request
req, err := ValidateSIOPRequest(requestURI)
if err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "invalid_request",
"error_description": err.Error(),
})
}
// Process as regular SIOP authorization
c.Set("siop_request", req)
return HandleSIOPAuthorization(c)
}
+353
View File
@@ -0,0 +1,353 @@
package handlers
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"github.com/labstack/echo/v4"
"github.com/stretchr/testify/assert"
)
// TestSIOPAuthorization tests Self-Issued OpenID Provider authorization
func TestSIOPAuthorization(t *testing.T) {
e := echo.New()
t.Run("ValidSIOPRequest", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/siop/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "id_token")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid did")
q.Set("nonce", "test-nonce")
q.Set("state", "test-state")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
// Set authenticated user DID
c.Set("user_did", "did:sonr:testuser")
err := HandleSIOPAuthorization(c)
assert.NoError(t, err)
// Should redirect with ID token
assert.Equal(t, http.StatusFound, rec.Code)
location := rec.Header().Get("Location")
assert.Contains(t, location, "id_token=")
assert.Contains(t, location, "state=test-state")
})
t.Run("InvalidResponseType", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/siop/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "code") // SIOP only supports id_token
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid")
q.Set("nonce", "test-nonce")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleSIOPAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
var errorResp map[string]string
err = json.Unmarshal(rec.Body.Bytes(), &errorResp)
assert.NoError(t, err)
assert.Equal(t, "unsupported_response_type", errorResp["error"])
})
t.Run("MissingScopeOpenID", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/siop/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "id_token")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "profile email") // Missing 'openid'
q.Set("nonce", "test-nonce")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleSIOPAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
var errorResp map[string]string
err = json.Unmarshal(rec.Body.Bytes(), &errorResp)
assert.NoError(t, err)
assert.Equal(t, "invalid_scope", errorResp["error"])
})
t.Run("UnauthenticatedUser", func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/siop/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "id_token")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid")
q.Set("nonce", "test-nonce")
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
// No user_did set - user not authenticated
err := HandleSIOPAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusUnauthorized, rec.Code)
var errorResp map[string]string
err = json.Unmarshal(rec.Body.Bytes(), &errorResp)
assert.NoError(t, err)
assert.Equal(t, "authentication_required", errorResp["error"])
})
}
// TestSIOPRegistration tests dynamic client registration for SIOP
func TestSIOPRegistration(t *testing.T) {
e := echo.New()
t.Run("ValidRegistration", func(t *testing.T) {
registration := map[string]any{
"client_name": "Test SIOP Client",
"client_purpose": "Testing SIOP functionality",
"logo_uri": "https://example.com/logo.png",
}
body, _ := json.Marshal(registration)
req := httptest.NewRequest(http.MethodPost, "/siop/register", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleSIOPRegistration(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var response map[string]any
err = json.Unmarshal(rec.Body.Bytes(), &response)
assert.NoError(t, err)
assert.NotEmpty(t, response["client_id"])
assert.Equal(t, "Test SIOP Client", response["client_name"])
assert.Contains(t, response["subject_syntax_types_supported"], "did")
})
t.Run("MissingClientName", func(t *testing.T) {
registration := map[string]any{
"client_purpose": "Testing",
}
body, _ := json.Marshal(registration)
req := httptest.NewRequest(http.MethodPost, "/siop/register", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleSIOPRegistration(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
var errorResp map[string]string
err = json.Unmarshal(rec.Body.Bytes(), &errorResp)
assert.NoError(t, err)
assert.Equal(t, "invalid_request", errorResp["error"])
assert.Contains(t, errorResp["error_description"], "Client name is required")
})
}
// TestSIOPMetadata tests SIOP metadata endpoint
func TestSIOPMetadata(t *testing.T) {
e := echo.New()
req := httptest.NewRequest(http.MethodGet, "/siop/metadata", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleSIOPMetadata(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusOK, rec.Code)
var metadata map[string]any
err = json.Unmarshal(rec.Body.Bytes(), &metadata)
assert.NoError(t, err)
// Verify required SIOP metadata fields
assert.Equal(t, "https://self-issued.me/v2", metadata["issuer"])
assert.Equal(t, "openid://", metadata["authorization_endpoint"])
assert.Contains(t, metadata["response_types_supported"], "id_token")
assert.Contains(t, metadata["scopes_supported"], "openid")
assert.Contains(t, metadata["subject_types_supported"], "pairwise")
assert.NotNil(t, metadata["vp_formats_supported"])
}
// TestValidateSIOPRequest tests SIOP request validation
func TestValidateSIOPRequest(t *testing.T) {
t.Run("ValidRequest", func(t *testing.T) {
requestURI := "openid://?" +
"response_type=id_token" +
"&client_id=test-client" +
"&redirect_uri=" + url.QueryEscape("http://localhost:3000/callback") +
"&scope=openid+did" +
"&nonce=test-nonce" +
"&state=test-state"
req, err := ValidateSIOPRequest(requestURI)
assert.NoError(t, err)
assert.NotNil(t, req)
assert.Equal(t, "id_token", req.ResponseType)
assert.Equal(t, "test-client", req.ClientID)
assert.Equal(t, "openid did", req.Scope)
assert.Equal(t, "test-nonce", req.Nonce)
})
t.Run("InvalidURI", func(t *testing.T) {
req, err := ValidateSIOPRequest("not-a-valid-uri")
assert.Error(t, err)
assert.Nil(t, req)
})
t.Run("WithClaims", func(t *testing.T) {
claims := map[string]any{
"id_token": map[string]any{
"email": map[string]any{
"essential": true,
},
"name": nil,
},
}
claimsJSON, _ := json.Marshal(claims)
requestURI := "openid://?" +
"response_type=id_token" +
"&client_id=test-client" +
"&redirect_uri=" + url.QueryEscape("http://localhost:3000/callback") +
"&scope=openid" +
"&nonce=test-nonce" +
"&claims=" + url.QueryEscape(string(claimsJSON))
req, err := ValidateSIOPRequest(requestURI)
assert.NoError(t, err)
assert.NotNil(t, req)
assert.NotNil(t, req.Claims.IDToken)
})
}
// TestSIOPWithVPToken tests SIOP with Verifiable Presentation
func TestSIOPWithVPToken(t *testing.T) {
e := echo.New()
vpClaims := map[string]any{
"vp_token": map[string]any{
"presentation_definition": map[string]any{
"id": "test-presentation",
"input_descriptors": []map[string]any{
{
"id": "id_credential",
"constraints": map[string]any{
"fields": []map[string]any{
{
"path": []string{"$.type"},
"filter": map[string]any{
"type": "string",
"const": "IdentityCredential",
},
},
},
},
},
},
},
},
}
claimsJSON, _ := json.Marshal(vpClaims)
req := httptest.NewRequest(http.MethodPost, "/siop/authorize", nil)
q := req.URL.Query()
q.Set("response_type", "id_token")
q.Set("client_id", "test-client")
q.Set("redirect_uri", "http://localhost:3000/callback")
q.Set("scope", "openid")
q.Set("nonce", "test-nonce")
q.Set("claims", string(claimsJSON))
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
c.Set("user_did", "did:sonr:testuser")
// Parse request to set claims
siopReq := &SIOPRequest{
ResponseType: "id_token",
ClientID: "test-client",
RedirectURI: "http://localhost:3000/callback",
Scope: "openid",
Nonce: "test-nonce",
}
err := json.Unmarshal(claimsJSON, &siopReq.Claims)
assert.NoError(t, err)
c.Set("siop_request", siopReq)
err = HandleSIOPAuthorization(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusFound, rec.Code)
location := rec.Header().Get("Location")
assert.Contains(t, location, "id_token=")
// VP token generation is optional, so we don't assert its presence
}
// TestHandleSIOPRequest tests complete SIOP request flow
func TestHandleSIOPRequest(t *testing.T) {
e := echo.New()
t.Run("ValidRequestURI", func(t *testing.T) {
requestURI := "openid://?" +
"response_type=id_token" +
"&client_id=test-client" +
"&redirect_uri=" + url.QueryEscape("http://localhost:3000/callback") +
"&scope=openid" +
"&nonce=test-nonce"
req := httptest.NewRequest(http.MethodGet, "/siop/request", nil)
q := req.URL.Query()
q.Set("request_uri", requestURI)
req.URL.RawQuery = q.Encode()
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
c.Set("user_did", "did:sonr:testuser")
err := HandleSIOPRequest(c)
assert.NoError(t, err)
// Should process as authorization request
assert.Equal(t, http.StatusFound, rec.Code)
})
t.Run("MissingRequestURI", func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "/siop/request", nil)
rec := httptest.NewRecorder()
c := e.NewContext(req, rec)
err := HandleSIOPRequest(c)
assert.NoError(t, err)
assert.Equal(t, http.StatusBadRequest, rec.Code)
var errorResp map[string]string
err = json.Unmarshal(rec.Body.Bytes(), &errorResp)
assert.NoError(t, err)
assert.Equal(t, "invalid_request", errorResp["error"])
assert.Contains(t, errorResp["error_description"], "Missing request_uri")
})
}
+265
View File
@@ -0,0 +1,265 @@
package handlers
import (
"time"
"github.com/golang-jwt/jwt/v5"
)
// WebAuthnCredential represents a WebAuthn credential
type WebAuthnCredential struct {
CredentialID string `json:"credential_id"`
RawID string `json:"raw_id"`
ClientDataJSON string `json:"client_data_json"`
AttestationObject string `json:"attestation_object"`
Username string `json:"username"`
Origin string `json:"origin"`
PublicKey []byte `json:"public_key"`
Algorithm int32 `json:"algorithm"`
CreatedAt time.Time `json:"created_at"`
}
// WebAuthnRegistrationRequest represents a registration request
type WebAuthnRegistrationRequest struct {
Username string `json:"username"`
Challenge string `json:"challenge,omitempty"`
AutoCreateVault bool `json:"auto_create_vault"`
BroadcastToChain bool `json:"broadcast_to_chain"`
}
// WebAuthnAuthenticationRequest represents an authentication request
type WebAuthnAuthenticationRequest struct {
Username string `json:"username"`
Challenge string `json:"challenge,omitempty"`
}
// WebAuthnRegistrationResponse contains registration ceremony data
type WebAuthnRegistrationResponse struct {
Challenge string `json:"challenge"`
RP WebAuthnRPEntity `json:"rp"`
User WebAuthnUserEntity `json:"user"`
PubKeyCredParams []WebAuthnCredParam `json:"pubKeyCredParams"`
AuthenticatorSelection WebAuthnAuthenticatorSelection `json:"authenticatorSelection"`
Timeout int `json:"timeout"`
Attestation string `json:"attestation"`
}
// WebAuthnAuthenticationResponse contains authentication ceremony data
type WebAuthnAuthenticationResponse struct {
Challenge string `json:"challenge"`
Timeout int `json:"timeout"`
RPID string `json:"rpId"`
AllowCredentials []WebAuthnAllowedCred `json:"allowCredentials"`
UserVerification string `json:"userVerification"`
}
// WebAuthnRPEntity represents relying party info
type WebAuthnRPEntity struct {
ID string `json:"id"`
Name string `json:"name"`
}
// WebAuthnUserEntity represents user info
type WebAuthnUserEntity struct {
ID string `json:"id"`
Name string `json:"name"`
DisplayName string `json:"displayName"`
}
// WebAuthnCredParam represents credential parameters
type WebAuthnCredParam struct {
Type string `json:"type"`
Alg int `json:"alg"`
}
// WebAuthnAuthenticatorSelection represents authenticator requirements
type WebAuthnAuthenticatorSelection struct {
AuthenticatorAttachment string `json:"authenticatorAttachment,omitempty"`
UserVerification string `json:"userVerification"`
ResidentKey string `json:"residentKey,omitempty"`
}
// WebAuthnAllowedCred represents allowed credentials for authentication
type WebAuthnAllowedCred struct {
Type string `json:"type"`
ID string `json:"id"`
}
// OIDCConfig represents OpenID Connect configuration
type OIDCConfig struct {
Issuer string `json:"issuer"`
AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"`
UserInfoEndpoint string `json:"userinfo_endpoint"`
JWKSEndpoint string `json:"jwks_uri"`
RevocationEndpoint string `json:"revocation_endpoint,omitempty"`
IntrospectionEndpoint string `json:"introspection_endpoint,omitempty"`
RegistrationEndpoint string `json:"registration_endpoint,omitempty"`
ScopesSupported []string `json:"scopes_supported"`
ResponseTypesSupported []string `json:"response_types_supported"`
GrantTypesSupported []string `json:"grant_types_supported"`
SubjectTypesSupported []string `json:"subject_types_supported"`
IDTokenSigningAlgValuesSupported []string `json:"id_token_signing_alg_values_supported"`
TokenEndpointAuthMethodsSupported []string `json:"token_endpoint_auth_methods_supported"`
ClaimsSupported []string `json:"claims_supported"`
CodeChallengeMethodsSupported []string `json:"code_challenge_methods_supported"`
}
// OIDCAuthorizationRequest represents an authorization request
type OIDCAuthorizationRequest struct {
ResponseType string `json:"response_type"`
ClientID string `json:"client_id"`
RedirectURI string `json:"redirect_uri"`
Scope string `json:"scope"`
State string `json:"state"`
Nonce string `json:"nonce,omitempty"`
CodeChallenge string `json:"code_challenge,omitempty"`
CodeChallengeMethod string `json:"code_challenge_method,omitempty"`
}
// OIDCTokenRequest represents a token request
type OIDCTokenRequest struct {
GrantType string `json:"grant_type"`
Code string `json:"code,omitempty"`
RedirectURI string `json:"redirect_uri,omitempty"`
ClientID string `json:"client_id"`
ClientSecret string `json:"client_secret,omitempty"`
CodeVerifier string `json:"code_verifier,omitempty"`
RefreshToken string `json:"refresh_token,omitempty"`
}
// OIDCTokenResponse represents a token response
type OIDCTokenResponse struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
RefreshToken string `json:"refresh_token,omitempty"`
IDToken string `json:"id_token,omitempty"`
Scope string `json:"scope,omitempty"`
}
// OIDCUserInfo represents user information
type OIDCUserInfo struct {
Subject string `json:"sub"`
Name string `json:"name,omitempty"`
PreferredUsername string `json:"preferred_username,omitempty"`
Email string `json:"email,omitempty"`
EmailVerified bool `json:"email_verified,omitempty"`
DID string `json:"did,omitempty"`
VaultID string `json:"vault_id,omitempty"`
UpdatedAt int64 `json:"updated_at,omitempty"`
Claims map[string]any `json:"claims,omitempty"`
}
// JWKSet represents a JSON Web Key Set
type JWKSet struct {
Keys []JWK `json:"keys"`
}
// JWK represents a JSON Web Key
type JWK struct {
KeyType string `json:"kty"`
Use string `json:"use,omitempty"`
KeyID string `json:"kid"`
Algorithm string `json:"alg,omitempty"`
N string `json:"n,omitempty"` // RSA modulus
E string `json:"e,omitempty"` // RSA exponent
X string `json:"x,omitempty"` // EC x coordinate
Y string `json:"y,omitempty"` // EC y coordinate
Curve string `json:"crv,omitempty"` // EC curve
}
// SIOPRequest represents a Self-Issued OpenID Provider request
type SIOPRequest struct {
ResponseType string `json:"response_type"`
ClientID string `json:"client_id"`
RedirectURI string `json:"redirect_uri"`
Scope string `json:"scope"`
Nonce string `json:"nonce"`
State string `json:"state,omitempty"`
Claims SIOPClaims `json:"claims,omitempty"`
Registration SIOPRegistration `json:"registration,omitempty"`
}
// SIOPClaims represents claims requested in SIOP
type SIOPClaims struct {
IDToken map[string]ClaimRequest `json:"id_token,omitempty"`
VPToken map[string]ClaimRequest `json:"vp_token,omitempty"`
}
// ClaimRequest represents a claim request
type ClaimRequest struct {
Essential bool `json:"essential,omitempty"`
Value string `json:"value,omitempty"`
Values []string `json:"values,omitempty"`
}
// SIOPRegistration represents client registration in SIOP
type SIOPRegistration struct {
ClientName string `json:"client_name,omitempty"`
ClientPurpose string `json:"client_purpose,omitempty"`
LogoURI string `json:"logo_uri,omitempty"`
SubjectSyntaxTypesSupported []string `json:"subject_syntax_types_supported,omitempty"`
VPFormats map[string]any `json:"vp_formats,omitempty"`
}
// SIOPResponse represents a Self-Issued OpenID Provider response
type SIOPResponse struct {
IDToken string `json:"id_token"`
VPToken string `json:"vp_token,omitempty"`
State string `json:"state,omitempty"`
}
// DIDAuthClaims represents DID-based authentication claims
type DIDAuthClaims struct {
jwt.RegisteredClaims
DID string `json:"did"`
Challenge string `json:"challenge,omitempty"`
WebAuthn *WebAuthnCredential `json:"webauthn,omitempty"`
Extra map[string]any `json:"extra,omitempty"`
}
// AuthorizationCode represents an OAuth2 authorization code
type AuthorizationCode struct {
Code string
ClientID string
RedirectURI string
UserDID string
Scope string
Nonce string
CodeChallenge string
CodeChallengeMethod string
ExpiresAt time.Time
Used bool
}
// OIDCSession represents an active OIDC session
type OIDCSession struct {
SessionID string
UserDID string
ClientID string
Scope string
Nonce string
AccessToken string
RefreshToken string
IDToken string
ExpiresAt time.Time
CreatedAt time.Time
}
// BroadcastRequest represents a blockchain broadcast request
type BroadcastRequest struct {
Message any `json:"message"`
Gasless bool `json:"gasless"`
AutoSign bool `json:"auto_sign"`
FromAddress string `json:"from_address,omitempty"`
}
// BroadcastResponse represents a blockchain broadcast response
type BroadcastResponse struct {
TxHash string `json:"tx_hash"`
Height int64 `json:"height,omitempty"`
Code uint32 `json:"code,omitempty"`
RawLog string `json:"raw_log,omitempty"`
Success bool `json:"success"`
}
+602
View File
@@ -0,0 +1,602 @@
package handlers
import (
"context"
"crypto/ed25519"
"crypto/rand"
"encoding/base64"
"encoding/json"
"fmt"
"time"
"github.com/cosmos/cosmos-sdk/crypto/keys/secp256k1"
"github.com/golang-jwt/jwt/v5"
"github.com/sonr-io/sonr/crypto/ucan"
"github.com/sonr-io/sonr/x/did/types"
)
// BlockchainUCANSigner implements UCAN signing with blockchain keys
type BlockchainUCANSigner struct {
didKeeper DIDKeeperInterface
issuerDID string
signingMethod jwt.SigningMethod
privateKey any
}
// DIDKeeperInterface defines the minimal interface needed from DID keeper
type DIDKeeperInterface interface {
GetDIDDocument(ctx context.Context, did string) (*types.DIDDocument, error)
GetVerificationMethod(
ctx context.Context,
did string,
methodID string,
) (*types.VerificationMethod, error)
}
// NewBlockchainUCANSigner creates a new blockchain-integrated UCAN signer
func NewBlockchainUCANSigner(
didKeeper DIDKeeperInterface,
issuerDID string,
) (*BlockchainUCANSigner, error) {
return &BlockchainUCANSigner{
didKeeper: didKeeper,
issuerDID: issuerDID,
signingMethod: jwt.SigningMethodES256, // Default to ES256
}, nil
}
// Sign signs a UCAN token with blockchain keys
func (s *BlockchainUCANSigner) Sign(token *ucan.Token) (string, error) {
// Build JWT claims from UCAN token
claims := s.buildClaims(token)
// Get signing key
signingKey, method, err := s.getSigningKey(context.Background())
if err != nil {
return "", fmt.Errorf("failed to get signing key: %w", err)
}
// Create JWT token
jwtToken := jwt.NewWithClaims(method, claims)
// Sign token
signedToken, err := jwtToken.SignedString(signingKey)
if err != nil {
return "", fmt.Errorf("failed to sign token: %w", err)
}
return signedToken, nil
}
// GetIssuerDID returns the issuer DID
func (s *BlockchainUCANSigner) GetIssuerDID() string {
if s.issuerDID == "" {
return "did:sonr:oauth-provider"
}
return s.issuerDID
}
// buildClaims builds JWT claims from UCAN token
func (s *BlockchainUCANSigner) buildClaims(token *ucan.Token) jwt.MapClaims {
claims := jwt.MapClaims{
"iss": token.Issuer,
"aud": token.Audience,
"exp": token.ExpiresAt,
"iat": time.Now().Unix(),
"nbf": token.NotBefore,
}
// Add attenuations
if len(token.Attenuations) > 0 {
attClaims := make([]map[string]any, len(token.Attenuations))
for i, att := range token.Attenuations {
attClaims[i] = s.serializeAttenuation(att)
}
claims["att"] = attClaims
}
// Add proofs if present
if len(token.Proofs) > 0 {
proofClaims := make([]string, len(token.Proofs))
for i, proof := range token.Proofs {
proofClaims[i] = string(proof)
}
claims["prf"] = proofClaims
}
// Add facts if present
if len(token.Facts) > 0 {
factClaims := make([]json.RawMessage, len(token.Facts))
for i, fact := range token.Facts {
factClaims[i] = fact.Data
}
claims["fct"] = factClaims
}
// Add UCAN version
claims["ucv"] = "0.10.0"
return claims
}
// serializeAttenuation serializes an attenuation for JWT claims
func (s *BlockchainUCANSigner) serializeAttenuation(att ucan.Attenuation) map[string]any {
result := map[string]any{
"with": att.Resource.GetURI(),
}
// Handle capability serialization
actions := att.Capability.GetActions()
if len(actions) == 1 {
result["can"] = actions[0]
} else {
result["can"] = actions
}
// Add resource-specific metadata
scheme := att.Resource.GetScheme()
switch scheme {
case "vault":
result["type"] = "vault_operation"
case "service", "svc":
result["type"] = "service_operation"
case "did":
result["type"] = "identity_operation"
case "dwn":
result["type"] = "data_operation"
case "dex", "pool":
result["type"] = "trading_operation"
}
return result
}
// getSigningKey retrieves the signing key from blockchain
func (s *BlockchainUCANSigner) getSigningKey(ctx context.Context) (any, jwt.SigningMethod, error) {
// If we have a cached private key, use it
if s.privateKey != nil {
return s.privateKey, s.signingMethod, nil
}
// For OAuth provider, use a service key
if s.issuerDID == "did:sonr:oauth-provider" || s.issuerDID == "" {
// Generate or retrieve service key
privateKey, publicKey, err := s.generateServiceKey()
if err != nil {
return nil, nil, err
}
s.privateKey = privateKey
s.signingMethod = jwt.SigningMethodEdDSA // Update to EdDSA for Ed25519 keys
_ = publicKey // Store public key if needed
return privateKey, jwt.SigningMethodEdDSA, nil
}
// For DID-based signing, retrieve from DID document
if s.didKeeper != nil {
didDoc, err := s.didKeeper.GetDIDDocument(ctx, s.issuerDID)
if err != nil {
return nil, nil, fmt.Errorf("failed to get DID document: %w", err)
}
// Use the first verification method for signing
if len(didDoc.VerificationMethod) > 0 {
vm := didDoc.VerificationMethod[0]
// Extract key based on verification method kind
switch vm.VerificationMethodKind {
case "Ed25519VerificationKey2020":
// Extract Ed25519 key
privateKey, err := s.extractEd25519Key(vm)
if err != nil {
return nil, nil, err
}
s.privateKey = privateKey
s.signingMethod = jwt.SigningMethodEdDSA
return privateKey, jwt.SigningMethodEdDSA, nil
case "EcdsaSecp256k1VerificationKey2019":
// Extract Secp256k1 key
privateKey, err := s.extractSecp256k1Key(vm)
if err != nil {
return nil, nil, err
}
s.privateKey = privateKey
s.signingMethod = jwt.SigningMethodES256
return privateKey, jwt.SigningMethodES256, nil
default:
return nil, nil, fmt.Errorf(
"unsupported verification method type: %s",
vm.VerificationMethodKind,
)
}
}
}
return nil, nil, fmt.Errorf("no signing key available")
}
// generateServiceKey generates a new service key pair
func (s *BlockchainUCANSigner) generateServiceKey() (ed25519.PrivateKey, ed25519.PublicKey, error) {
// In production, this should retrieve from secure storage
// For now, generate a new key pair
publicKey, privateKey, err := ed25519.GenerateKey(rand.Reader)
if err != nil {
return nil, nil, fmt.Errorf("failed to generate key pair: %w", err)
}
return privateKey, publicKey, nil
}
// extractEd25519Key extracts Ed25519 private key from verification method
func (s *BlockchainUCANSigner) extractEd25519Key(
vm *types.VerificationMethod,
) (ed25519.PrivateKey, error) {
// In production, this would retrieve the private key from secure storage
// based on the public key in the verification method
// For now, return error as we don't have access to private keys
return nil, fmt.Errorf("private key retrieval not implemented")
}
// extractSecp256k1Key extracts Secp256k1 private key from verification method
func (s *BlockchainUCANSigner) extractSecp256k1Key(
vm *types.VerificationMethod,
) (*secp256k1.PrivKey, error) {
// In production, this would retrieve the private key from secure storage
// based on the public key in the verification method
// For now, return error as we don't have access to private keys
return nil, fmt.Errorf("private key retrieval not implemented")
}
// VerifySignature verifies a UCAN token signature
func (s *BlockchainUCANSigner) VerifySignature(tokenString string) (*ucan.Token, error) {
// Parse token without verification first to get issuer
token, _, err := new(jwt.Parser).ParseUnverified(tokenString, jwt.MapClaims{})
if err != nil {
return nil, fmt.Errorf("failed to parse token: %w", err)
}
claims, ok := token.Claims.(jwt.MapClaims)
if !ok {
return nil, fmt.Errorf("invalid claims format")
}
issuer, ok := claims["iss"].(string)
if !ok {
return nil, fmt.Errorf("no issuer in token")
}
// Get public key for issuer
publicKey, err := s.getPublicKey(context.Background(), issuer)
if err != nil {
return nil, fmt.Errorf("failed to get public key: %w", err)
}
// Verify token with public key
parsedToken, err := jwt.Parse(tokenString, func(token *jwt.Token) (any, error) {
return publicKey, nil
})
if err != nil {
return nil, fmt.Errorf("token verification failed: %w", err)
}
if !parsedToken.Valid {
return nil, fmt.Errorf("invalid token")
}
// Convert JWT claims to UCAN token
return s.claimsToUCAN(claims, tokenString)
}
// getPublicKey retrieves public key for a DID
func (s *BlockchainUCANSigner) getPublicKey(ctx context.Context, did string) (any, error) {
if did == "did:sonr:oauth-provider" {
// Return service public key
// In production, this would be retrieved from configuration
return []byte("oauth-provider-public-key"), nil
}
if s.didKeeper != nil {
didDoc, err := s.didKeeper.GetDIDDocument(ctx, did)
if err != nil {
return nil, fmt.Errorf("failed to get DID document: %w", err)
}
if len(didDoc.VerificationMethod) > 0 {
vm := didDoc.VerificationMethod[0]
// Extract public key based on verification method kind
switch vm.VerificationMethodKind {
case "Ed25519VerificationKey2020":
// Decode base64 public key
publicKeyBytes, err := base64.StdEncoding.DecodeString(vm.PublicKeyMultibase)
if err != nil {
return nil, err
}
return ed25519.PublicKey(publicKeyBytes), nil
case "EcdsaSecp256k1VerificationKey2019":
// Decode and return Secp256k1 public key
publicKeyBytes, err := base64.StdEncoding.DecodeString(vm.PublicKeyMultibase)
if err != nil {
return nil, err
}
return publicKeyBytes, nil
default:
return nil, fmt.Errorf(
"unsupported verification method type: %s",
vm.VerificationMethodKind,
)
}
}
}
return nil, fmt.Errorf("no public key found for DID: %s", did)
}
// claimsToUCAN converts JWT claims to UCAN token
func (s *BlockchainUCANSigner) claimsToUCAN(
claims jwt.MapClaims,
rawToken string,
) (*ucan.Token, error) {
token := &ucan.Token{
Raw: rawToken,
}
// Extract standard claims
if iss, ok := claims["iss"].(string); ok {
token.Issuer = iss
}
if aud, ok := claims["aud"].(string); ok {
token.Audience = aud
}
if exp, ok := claims["exp"].(float64); ok {
token.ExpiresAt = int64(exp)
}
if nbf, ok := claims["nbf"].(float64); ok {
token.NotBefore = int64(nbf)
}
// Extract attenuations
if attClaims, ok := claims["att"].([]any); ok {
attenuations := make([]ucan.Attenuation, 0, len(attClaims))
for _, attItem := range attClaims {
if attMap, ok := attItem.(map[string]any); ok {
att, err := s.parseAttenuation(attMap)
if err != nil {
return nil, fmt.Errorf("failed to parse attenuation: %w", err)
}
attenuations = append(attenuations, att)
}
}
token.Attenuations = attenuations
}
// Extract proofs
if proofClaims, ok := claims["prf"].([]any); ok {
proofs := make([]ucan.Proof, 0, len(proofClaims))
for _, proofItem := range proofClaims {
if proofStr, ok := proofItem.(string); ok {
proofs = append(proofs, ucan.Proof(proofStr))
}
}
token.Proofs = proofs
}
// Extract facts
if factClaims, ok := claims["fct"].([]any); ok {
facts := make([]ucan.Fact, 0, len(factClaims))
for _, factItem := range factClaims {
factData, _ := json.Marshal(factItem)
facts = append(facts, ucan.Fact{
Data: json.RawMessage(factData),
})
}
token.Facts = facts
}
return token, nil
}
// parseAttenuation parses an attenuation from JWT claims
func (s *BlockchainUCANSigner) parseAttenuation(attMap map[string]any) (ucan.Attenuation, error) {
// Extract resource URI
resourceURI, ok := attMap["with"].(string)
if !ok {
return ucan.Attenuation{}, fmt.Errorf("missing resource URI")
}
// Parse resource
resource := &SimpleResource{
Scheme: "generic",
Value: resourceURI,
}
// Extract scheme from URI if possible
if len(resourceURI) > 0 {
for _, scheme := range []string{"vault:", "service:", "did:", "dwn:", "dex:", "pool:"} {
if len(resourceURI) >= len(scheme) && resourceURI[:len(scheme)] == scheme {
resource.Scheme = scheme[:len(scheme)-1]
resource.Value = resourceURI[len(scheme):]
break
}
}
}
// Extract capability
var capability ucan.Capability
switch can := attMap["can"].(type) {
case string:
capability = &ucan.SimpleCapability{Action: can}
case []any:
actions := make([]string, 0, len(can))
for _, action := range can {
if actionStr, ok := action.(string); ok {
actions = append(actions, actionStr)
}
}
capability = &ucan.MultiCapability{Actions: actions}
default:
return ucan.Attenuation{}, fmt.Errorf("invalid capability format")
}
return ucan.Attenuation{
Capability: capability,
Resource: resource,
}, nil
}
// SetPrivateKey sets a private key for signing (for testing)
func (s *BlockchainUCANSigner) SetPrivateKey(privateKey any, method jwt.SigningMethod) {
s.privateKey = privateKey
s.signingMethod = method
}
// CreateDelegationToken creates a UCAN token for delegation
func (s *BlockchainUCANSigner) CreateDelegationToken(
issuer, audience string,
attenuations []ucan.Attenuation,
proofs []ucan.Proof,
expiresIn time.Duration,
) (string, error) {
token := &ucan.Token{
Issuer: issuer,
Audience: audience,
ExpiresAt: time.Now().Add(expiresIn).Unix(),
NotBefore: time.Now().Unix(),
Attenuations: attenuations,
Proofs: proofs,
}
return s.Sign(token)
}
// ValidateDelegationChain validates a chain of UCAN delegations
func (s *BlockchainUCANSigner) ValidateDelegationChain(tokens []string) error {
if len(tokens) == 0 {
return fmt.Errorf("empty delegation chain")
}
var previousToken *ucan.Token
for i, tokenString := range tokens {
// Verify current token
token, err := s.VerifySignature(tokenString)
if err != nil {
return fmt.Errorf("failed to verify token %d: %w", i, err)
}
// Check expiration
if time.Now().Unix() > token.ExpiresAt {
return fmt.Errorf("token %d has expired", i)
}
// Check not before
if token.NotBefore > 0 && time.Now().Unix() < token.NotBefore {
return fmt.Errorf("token %d not yet valid", i)
}
// Validate delegation chain
if previousToken != nil {
// Check that previous token's audience matches current issuer
if previousToken.Audience != token.Issuer {
return fmt.Errorf("broken delegation chain at token %d", i)
}
// Check that capabilities are properly attenuated
if !s.isProperlyAttenuated(previousToken.Attenuations, token.Attenuations) {
return fmt.Errorf("improper attenuation at token %d", i)
}
}
previousToken = token
}
return nil
}
// isProperlyAttenuated checks if child attenuations are properly attenuated from parent
func (s *BlockchainUCANSigner) isProperlyAttenuated(parent, child []ucan.Attenuation) bool {
// Child cannot have more permissions than parent
for _, childAtt := range child {
found := false
for _, parentAtt := range parent {
// Check if child attenuation is covered by parent
if s.attenuationCovers(parentAtt, childAtt) {
found = true
break
}
}
if !found {
return false
}
}
return true
}
// attenuationCovers checks if parent attenuation covers child
func (s *BlockchainUCANSigner) attenuationCovers(parent, child ucan.Attenuation) bool {
// Check resource match
if parent.Resource.GetScheme() != child.Resource.GetScheme() {
// Wildcard scheme matches all
if parent.Resource.GetScheme() != "*" {
return false
}
}
// Check capability coverage
parentActions := parent.Capability.GetActions()
childActions := child.Capability.GetActions()
// If parent has wildcard, it covers everything
for _, action := range parentActions {
if action == "*" {
return true
}
}
// Check each child action is in parent
for _, childAction := range childActions {
found := false
for _, parentAction := range parentActions {
if parentAction == childAction {
found = true
break
}
}
if !found {
return false
}
}
return true
}
// RefreshToken creates a new token from an existing one with updated expiration
func (s *BlockchainUCANSigner) RefreshToken(
tokenString string,
newExpiration time.Duration,
) (string, error) {
// Verify existing token
token, err := s.VerifySignature(tokenString)
if err != nil {
return "", fmt.Errorf("failed to verify token: %w", err)
}
// Create new token with same claims but new expiration
newToken := &ucan.Token{
Issuer: token.Issuer,
Audience: token.Audience,
ExpiresAt: time.Now().Add(newExpiration).Unix(),
NotBefore: time.Now().Unix(),
Attenuations: token.Attenuations,
Proofs: append(token.Proofs, ucan.Proof(tokenString)), // Add original as proof
Facts: token.Facts,
}
return s.Sign(newToken)
}
+535
View File
@@ -0,0 +1,535 @@
package handlers
import (
"fmt"
"net/http"
"time"
"github.com/golang-jwt/jwt/v5"
"github.com/hibiken/asynq"
"github.com/labstack/echo/v4"
"github.com/sonr-io/sonr/bridge/tasks"
"github.com/sonr-io/sonr/crypto/mpc"
"github.com/sonr-io/sonr/types/ipfs"
)
// VaultHandlers holds all vault-related handlers and their dependencies
type VaultHandlers struct {
IPFSClient ipfs.IPFSClient
ConnectionManager *ConnectionManager
SSEManager *SSEManager
}
// NewVaultHandlers creates a new VaultHandlers instance
func NewVaultHandlers(
ipfsClient ipfs.IPFSClient,
connManager *ConnectionManager,
sseManager *SSEManager,
) *VaultHandlers {
return &VaultHandlers{
IPFSClient: ipfsClient,
ConnectionManager: connManager,
SSEManager: sseManager,
}
}
// GetQueueFromPriority determines the appropriate queue based on priority
func GetQueueFromPriority(priority string) string {
switch priority {
case "critical", "high":
return "critical"
case "low":
return "low"
default:
return "default"
}
}
// GenerateHandler handles vault generation requests
func (vh *VaultHandlers) GenerateHandler(client *asynq.Client) echo.HandlerFunc {
return func(c echo.Context) error {
// Extract user info from JWT token
user := c.Get("user").(*jwt.Token)
claims := user.Claims.(jwt.MapClaims)
userID := claims["user_id"].(string)
var req struct {
UserID int `json:"user_id"`
Priority string `json:"priority,omitempty"`
}
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
task, err := tasks.NewUCANDIDTask(req.UserID)
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to create task"},
)
}
queue := GetQueueFromPriority(req.Priority)
info, err := client.Enqueue(task, asynq.Queue(queue))
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to enqueue task"},
)
}
// Broadcast task status to WebSocket and SSE connections
go func() {
message := TaskStatusMessage{
TaskID: info.ID,
Status: "enqueued",
Time: time.Now(),
}
vh.ConnectionManager.BroadcastToTask(info.ID, message)
vh.SSEManager.BroadcastToSSE(info.ID, message)
// Simulate task progression for demonstration
// In a real implementation, this would be triggered by the actual task processors
time.Sleep(1 * time.Second)
processingMessage := TaskStatusMessage{
TaskID: info.ID,
Status: "processing",
Progress: 50,
Time: time.Now(),
}
vh.ConnectionManager.BroadcastToTask(info.ID, processingMessage)
vh.SSEManager.BroadcastToSSE(info.ID, processingMessage)
// Simulate completion
time.Sleep(2 * time.Second)
completedMessage := TaskStatusMessage{
TaskID: info.ID,
Status: "completed",
Progress: 100,
Time: time.Now(),
}
vh.ConnectionManager.BroadcastToTask(info.ID, completedMessage)
vh.SSEManager.BroadcastToSSE(info.ID, completedMessage)
}()
return c.JSON(http.StatusOK, map[string]any{
"task_id": info.ID,
"queue": queue,
"status": "enqueued",
"user_id": userID,
})
}
}
// SignHandler handles vault signing requests
func (vh *VaultHandlers) SignHandler(client *asynq.Client) echo.HandlerFunc {
return func(c echo.Context) error {
var req struct {
Message []byte `json:"message"`
Enclave *mpc.EnclaveData `json:"enclave,omitempty"` // Direct enclave data (fallback)
EnclaveCID string `json:"enclave_cid,omitempty"` // IPFS CID reference
Password []byte `json:"password,omitempty"` // Decryption password for IPFS stored data
UCANToken string `json:"ucan_token,omitempty"` // UCAN token for authorization
Priority string `json:"priority,omitempty"`
}
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
if len(req.Message) == 0 {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Message is required"})
}
// Handle enclave data - either direct or via IPFS CID
var enclave *mpc.EnclaveData
if req.EnclaveCID != "" {
// Retrieve enclave from IPFS
if vh.IPFSClient == nil {
return c.JSON(
http.StatusServiceUnavailable,
map[string]string{"error": "IPFS client not available"},
)
}
encryptedData, err := vh.IPFSClient.Get(req.EnclaveCID)
if err != nil {
return c.JSON(
http.StatusNotFound,
map[string]string{"error": "Failed to retrieve enclave from IPFS"},
)
}
// Create temporary enclave for decryption
tempEnclave := &mpc.EnclaveData{}
decryptedData, err := tempEnclave.Decrypt(req.Password, encryptedData)
if err != nil {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Failed to decrypt enclave data"},
)
}
// Unmarshal decrypted data
enclave = &mpc.EnclaveData{}
if err := enclave.Unmarshal(decryptedData); err != nil {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Failed to parse enclave data"},
)
}
} else if req.Enclave != nil {
// Use directly provided enclave data
enclave = req.Enclave
} else {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Either enclave or enclave_cid is required"})
}
task, err := tasks.NewUCANSignTask(
0,
req.Message,
) // TODO: Extract userID from JWT or context
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to create task"},
)
}
queue := GetQueueFromPriority(req.Priority)
info, err := client.Enqueue(task, asynq.Queue(queue))
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to enqueue task"},
)
}
return c.JSON(http.StatusOK, map[string]any{
"task_id": info.ID,
"queue": queue,
"status": "enqueued",
})
}
}
// VerifyHandler handles vault verification requests
func (vh *VaultHandlers) VerifyHandler(client *asynq.Client) echo.HandlerFunc {
return func(c echo.Context) error {
var req struct {
PublicKey []byte `json:"public_key"`
Message []byte `json:"message"`
Signature []byte `json:"signature"`
Priority string `json:"priority,omitempty"`
}
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
if len(req.PublicKey) == 0 || len(req.Message) == 0 || len(req.Signature) == 0 {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "PublicKey, message, and signature are required"},
)
}
task, err := tasks.NewUCANVerifyTask(
0,
req.Message,
req.Signature,
) // TODO: Extract userID from JWT or context
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to create task"},
)
}
queue := GetQueueFromPriority(req.Priority)
info, err := client.Enqueue(task, asynq.Queue(queue))
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to enqueue task"},
)
}
return c.JSON(http.StatusOK, map[string]any{
"task_id": info.ID,
"queue": queue,
"status": "enqueued",
})
}
}
// ExportHandler handles vault export requests
func (vh *VaultHandlers) ExportHandler(client *asynq.Client) echo.HandlerFunc {
return func(c echo.Context) error {
var req struct {
Enclave *mpc.EnclaveData `json:"enclave,omitempty"` // Direct enclave data
EnclaveCID string `json:"enclave_cid,omitempty"` // IPFS CID reference
Password []byte `json:"password"` // For encryption/decryption
StoreIPFS bool `json:"store_ipfs,omitempty"` // Whether to store result in IPFS
Priority string `json:"priority,omitempty"`
}
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
if len(req.Password) == 0 {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Password is required"})
}
// Handle enclave data - either direct or via IPFS CID
var enclave *mpc.EnclaveData
if req.EnclaveCID != "" {
// Retrieve enclave from IPFS
if vh.IPFSClient == nil {
return c.JSON(
http.StatusServiceUnavailable,
map[string]string{"error": "IPFS client not available"},
)
}
encryptedData, err := vh.IPFSClient.Get(req.EnclaveCID)
if err != nil {
return c.JSON(
http.StatusNotFound,
map[string]string{"error": "Failed to retrieve enclave from IPFS"},
)
}
// Create temporary enclave for decryption
tempEnclave := &mpc.EnclaveData{}
decryptedData, err := tempEnclave.Decrypt(req.Password, encryptedData)
if err != nil {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Failed to decrypt enclave data"},
)
}
// Unmarshal decrypted data
enclave = &mpc.EnclaveData{}
if err := enclave.Unmarshal(decryptedData); err != nil {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Failed to parse enclave data"},
)
}
} else if req.Enclave != nil {
// Use directly provided enclave data
enclave = req.Enclave
} else {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Either enclave or enclave_cid is required"})
}
// If store_ipfs is true, encrypt and store the enclave in IPFS first
if req.StoreIPFS && vh.IPFSClient != nil {
encryptedData, err := enclave.Encrypt(req.Password)
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to encrypt enclave data"},
)
}
cid, err := vh.IPFSClient.Add(encryptedData)
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to store enclave in IPFS"},
)
}
// Return the CID for future reference
return c.JSON(http.StatusOK, map[string]any{
"cid": cid,
"status": "stored",
"message": "Enclave data encrypted and stored in IPFS",
})
}
// Export functionality replaced with UCAN token creation for data access
// Convert export operation to UCAN token generation with export permissions
attenuations := []map[string]any{
{
"can": []string{"export", "read"},
"with": "vault://exported-data",
},
}
task, err := tasks.NewUCANTokenTask(
0,
"did:sonr:export-recipient",
attenuations,
time.Now().Add(24*time.Hour).Unix(),
)
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to create task"},
)
}
queue := GetQueueFromPriority(req.Priority)
info, err := client.Enqueue(task, asynq.Queue(queue))
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to enqueue task"},
)
}
return c.JSON(http.StatusOK, map[string]any{
"task_id": info.ID,
"queue": queue,
"status": "enqueued",
})
}
}
// ImportHandler handles vault import requests
func (vh *VaultHandlers) ImportHandler(client *asynq.Client) echo.HandlerFunc {
return func(c echo.Context) error {
var req struct {
CID string `json:"cid"`
Password []byte `json:"password"`
Priority string `json:"priority,omitempty"`
}
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
if req.CID == "" || len(req.Password) == 0 {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "CID and password are required"},
)
}
// Import functionality replaced with UCAN token creation for data import
// Convert import operation to UCAN token generation with import permissions
attenuations := []map[string]any{
{
"can": []string{"import", "write"},
"with": fmt.Sprintf("ipfs://%s", req.CID),
},
}
task, err := tasks.NewUCANTokenTask(
0,
"did:sonr:import-recipient",
attenuations,
time.Now().Add(1*time.Hour).Unix(),
)
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to create task"},
)
}
queue := GetQueueFromPriority(req.Priority)
info, err := client.Enqueue(task, asynq.Queue(queue))
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to enqueue task"},
)
}
return c.JSON(http.StatusOK, map[string]any{
"task_id": info.ID,
"queue": queue,
"status": "enqueued",
})
}
}
// RefreshHandler handles vault refresh requests
func (vh *VaultHandlers) RefreshHandler(client *asynq.Client) echo.HandlerFunc {
return func(c echo.Context) error {
var req struct {
Enclave *mpc.EnclaveData `json:"enclave,omitempty"` // Direct enclave data
EnclaveCID string `json:"enclave_cid,omitempty"` // IPFS CID reference
Password []byte `json:"password,omitempty"` // Decryption password for IPFS stored data
Priority string `json:"priority,omitempty"`
}
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Invalid JSON payload"})
}
// Handle enclave data - either direct or via IPFS CID
var enclave *mpc.EnclaveData
if req.EnclaveCID != "" {
// Retrieve enclave from IPFS
if vh.IPFSClient == nil {
return c.JSON(
http.StatusServiceUnavailable,
map[string]string{"error": "IPFS client not available"},
)
}
encryptedData, err := vh.IPFSClient.Get(req.EnclaveCID)
if err != nil {
return c.JSON(
http.StatusNotFound,
map[string]string{"error": "Failed to retrieve enclave from IPFS"},
)
}
// Create temporary enclave for decryption
tempEnclave := &mpc.EnclaveData{}
decryptedData, err := tempEnclave.Decrypt(req.Password, encryptedData)
if err != nil {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Failed to decrypt enclave data"},
)
}
// Unmarshal decrypted data
enclave = &mpc.EnclaveData{}
if err := enclave.Unmarshal(decryptedData); err != nil {
return c.JSON(
http.StatusBadRequest,
map[string]string{"error": "Failed to parse enclave data"},
)
}
} else if req.Enclave != nil {
// Use directly provided enclave data
enclave = req.Enclave
} else {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Either enclave or enclave_cid is required"})
}
// Refresh functionality replaced with UCAN DID generation for new identity
// Convert refresh operation to DID generation which includes key refresh
task, err := tasks.NewUCANDIDTask(0) // TODO: Extract userID from JWT or context
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to create task"},
)
}
queue := GetQueueFromPriority(req.Priority)
info, err := client.Enqueue(task, asynq.Queue(queue))
if err != nil {
return c.JSON(
http.StatusInternalServerError,
map[string]string{"error": "Failed to enqueue task"},
)
}
return c.JSON(http.StatusOK, map[string]any{
"task_id": info.ID,
"queue": queue,
"status": "enqueued",
})
}
}
+558
View File
@@ -0,0 +1,558 @@
package handlers
import (
"crypto/rand"
"encoding/base64"
"fmt"
"net/http"
"sync"
"time"
"github.com/labstack/echo/v4"
"github.com/sonr-io/sonr/types/webauthn"
"github.com/sonr-io/sonr/types/webauthn/webauthncbor"
)
// WebAuthnStore manages WebAuthn sessions and credentials
type WebAuthnStore struct {
mu sync.RWMutex
sessions map[string]*WebAuthnSession
credentials map[string][]*WebAuthnCredential
}
// WebAuthnSession holds session data for WebAuthn ceremonies
type WebAuthnSession struct {
Challenge string
Username string
CreatedAt time.Time
SessionType string // "registration" or "authentication"
}
var (
webAuthnStore = &WebAuthnStore{
sessions: make(map[string]*WebAuthnSession),
credentials: make(map[string][]*WebAuthnCredential),
}
sessionTimeout = 5 * time.Minute
)
// BeginWebAuthnRegistration starts WebAuthn registration ceremony
func BeginWebAuthnRegistration(c echo.Context) error {
var req WebAuthnRegistrationRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Invalid request format",
})
}
if req.Username == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Username is required",
})
}
// Store the request in context for later use in FinishWebAuthnRegistration
c.Set("webauthn_registration_request", &req)
// Generate challenge
challenge, err := generateChallenge()
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "Failed to generate challenge",
})
}
// Store session
session := &WebAuthnSession{
Challenge: challenge,
Username: req.Username,
CreatedAt: time.Now(),
SessionType: "registration",
}
webAuthnStore.mu.Lock()
webAuthnStore.sessions[req.Username] = session
webAuthnStore.mu.Unlock()
// Create registration response
response := WebAuthnRegistrationResponse{
Challenge: challenge,
RP: WebAuthnRPEntity{
ID: "localhost", // TODO: Get from config
Name: "Sonr Identity Platform",
},
User: WebAuthnUserEntity{
ID: base64.URLEncoding.EncodeToString([]byte(req.Username)),
Name: req.Username,
DisplayName: req.Username,
},
PubKeyCredParams: []WebAuthnCredParam{
{Type: "public-key", Alg: -7}, // ES256
{Type: "public-key", Alg: -257}, // RS256
},
AuthenticatorSelection: WebAuthnAuthenticatorSelection{
AuthenticatorAttachment: "platform",
UserVerification: "required",
ResidentKey: "preferred",
},
Timeout: 60000,
Attestation: "direct",
}
return c.JSON(http.StatusOK, response)
}
// FinishWebAuthnRegistration completes WebAuthn registration ceremony
func FinishWebAuthnRegistration(c echo.Context) error {
username := c.QueryParam("username")
if username == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Username is required",
})
}
// Parse registration response
var regResponse map[string]any
if err := c.Bind(&regResponse); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Invalid registration response",
})
}
// Get stored session
webAuthnStore.mu.RLock()
session, exists := webAuthnStore.sessions[username]
webAuthnStore.mu.RUnlock()
if !exists || session.SessionType != "registration" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "No registration session found",
})
}
// Check session timeout
if time.Since(session.CreatedAt) > sessionTimeout {
webAuthnStore.mu.Lock()
delete(webAuthnStore.sessions, username)
webAuthnStore.mu.Unlock()
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Registration session expired",
})
}
// Extract credential data
credentialID, _ := regResponse["id"].(string)
rawID, _ := regResponse["rawId"].(string)
response, _ := regResponse["response"].(map[string]any)
clientDataJSON, _ := response["clientDataJSON"].(string)
attestationObject, _ := response["attestationObject"].(string)
// Verify client data
if err := verifyClientData(clientDataJSON, session.Challenge, "webauthn.create"); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": fmt.Sprintf("Client data verification failed: %v", err),
})
}
// Extract public key from attestation
publicKey, algorithm, err := extractPublicKeyFromAttestation(attestationObject)
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": fmt.Sprintf("Failed to extract public key: %v", err),
})
}
// Create credential
credential := &WebAuthnCredential{
CredentialID: credentialID,
RawID: rawID,
ClientDataJSON: clientDataJSON,
AttestationObject: attestationObject,
Username: username,
Origin: "localhost", // TODO: Extract from client data
PublicKey: publicKey,
Algorithm: algorithm,
CreatedAt: time.Now(),
}
// Store credential
webAuthnStore.mu.Lock()
webAuthnStore.credentials[username] = append(webAuthnStore.credentials[username], credential)
delete(webAuthnStore.sessions, username)
webAuthnStore.mu.Unlock()
// Check if we should broadcast to blockchain
broadcastReq, ok := c.Get("broadcast_to_chain").(bool)
if !ok {
// Check from original request stored in context
if origReq, exists := c.Get("webauthn_registration_request").(*WebAuthnRegistrationRequest); exists {
broadcastReq = origReq.BroadcastToChain
}
}
var broadcastResult *BroadcastResponse
if broadcastReq {
// Broadcast WebAuthn credential to blockchain as gasless transaction
result, err := BroadcastWebAuthnRegistration(credential, true)
if err != nil {
// Log error but don't fail registration
c.Logger().Error("Failed to broadcast WebAuthn credential:", err)
} else {
broadcastResult = result
}
}
// Check if we should create a vault
autoCreateVault, ok := c.Get("auto_create_vault").(bool)
if !ok {
// Check from original request
if origReq, exists := c.Get("webauthn_registration_request").(*WebAuthnRegistrationRequest); exists {
autoCreateVault = origReq.AutoCreateVault
}
}
var vaultResult *BroadcastResponse
if autoCreateVault {
// Create vault for the user
vaultConfig := map[string]any{
"type": "standard",
"encryption": "AES256",
"owner": username,
}
userDID := fmt.Sprintf("did:sonr:%s", username)
result, err := BroadcastVaultCreation(userDID, vaultConfig)
if err != nil {
// Log error but don't fail registration
c.Logger().Error("Failed to create vault:", err)
} else {
vaultResult = result
}
}
finalResponse := map[string]any{
"success": true,
"message": "Registration completed successfully",
"credentialId": credentialID,
}
if broadcastResult != nil {
finalResponse["broadcast"] = broadcastResult
}
if vaultResult != nil {
finalResponse["vault"] = vaultResult
}
return c.JSON(http.StatusOK, finalResponse)
}
// BeginWebAuthnAuthentication starts WebAuthn authentication ceremony
func BeginWebAuthnAuthentication(c echo.Context) error {
var req WebAuthnAuthenticationRequest
if err := c.Bind(&req); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Invalid request format",
})
}
if req.Username == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Username is required",
})
}
// Check if user has credentials
webAuthnStore.mu.RLock()
credentials, exists := webAuthnStore.credentials[req.Username]
webAuthnStore.mu.RUnlock()
if !exists || len(credentials) == 0 {
return c.JSON(http.StatusNotFound, map[string]string{
"error": "No credentials found for user",
})
}
// Generate challenge
challenge, err := generateChallenge()
if err != nil {
return c.JSON(http.StatusInternalServerError, map[string]string{
"error": "Failed to generate challenge",
})
}
// Store session
session := &WebAuthnSession{
Challenge: challenge,
Username: req.Username,
CreatedAt: time.Now(),
SessionType: "authentication",
}
webAuthnStore.mu.Lock()
webAuthnStore.sessions[req.Username] = session
webAuthnStore.mu.Unlock()
// Build allowed credentials
allowCredentials := make([]WebAuthnAllowedCred, len(credentials))
for i, cred := range credentials {
allowCredentials[i] = WebAuthnAllowedCred{
Type: "public-key",
ID: cred.CredentialID,
}
}
// Create authentication response
response := WebAuthnAuthenticationResponse{
Challenge: challenge,
Timeout: 60000,
RPID: "localhost", // TODO: Get from config
AllowCredentials: allowCredentials,
UserVerification: "required",
}
return c.JSON(http.StatusOK, response)
}
// FinishWebAuthnAuthentication completes WebAuthn authentication ceremony
func FinishWebAuthnAuthentication(c echo.Context) error {
username := c.QueryParam("username")
if username == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Username is required",
})
}
// Parse authentication response
var authResponse map[string]any
if err := c.Bind(&authResponse); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Invalid authentication response",
})
}
// Get stored session
webAuthnStore.mu.RLock()
session, exists := webAuthnStore.sessions[username]
webAuthnStore.mu.RUnlock()
if !exists || session.SessionType != "authentication" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "No authentication session found",
})
}
// Check session timeout
if time.Since(session.CreatedAt) > sessionTimeout {
webAuthnStore.mu.Lock()
delete(webAuthnStore.sessions, username)
webAuthnStore.mu.Unlock()
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Authentication session expired",
})
}
// Extract response data
credentialID, _ := authResponse["id"].(string)
response, _ := authResponse["response"].(map[string]any)
clientDataJSON, _ := response["clientDataJSON"].(string)
authenticatorData, _ := response["authenticatorData"].(string)
signature, _ := response["signature"].(string)
userHandle, _ := response["userHandle"].(string)
// Verify client data
if err := verifyClientData(clientDataJSON, session.Challenge, "webauthn.get"); err != nil {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": fmt.Sprintf("Client data verification failed: %v", err),
})
}
// Find matching credential
webAuthnStore.mu.RLock()
credentials := webAuthnStore.credentials[username]
webAuthnStore.mu.RUnlock()
var matchedCredential *WebAuthnCredential
for _, cred := range credentials {
if cred.CredentialID == credentialID {
matchedCredential = cred
break
}
}
if matchedCredential == nil {
return c.JSON(http.StatusUnauthorized, map[string]string{
"error": "Invalid credential",
})
}
// TODO: Verify signature using the stored public key
// This would require implementing proper WebAuthn signature verification
// Clean up session
webAuthnStore.mu.Lock()
delete(webAuthnStore.sessions, username)
webAuthnStore.mu.Unlock()
// Create authenticated session
userDID := fmt.Sprintf("did:sonr:%s", username)
authSession := &OIDCSession{
SessionID: generateSessionID(),
UserDID: userDID,
ClientID: "webauthn-client",
Scope: "openid profile did vault",
AccessToken: generateAccessToken(userDID),
RefreshToken: generateRefreshToken(),
ExpiresAt: time.Now().Add(time.Hour),
CreatedAt: time.Now(),
}
// Store session for OIDC compatibility
oidcProvider.mu.Lock()
oidcProvider.sessions[authSession.AccessToken] = authSession
oidcProvider.mu.Unlock()
// Set user context for downstream handlers
c.Set("user_did", userDID)
c.Set("authenticated", true)
c.Set("auth_method", "webauthn")
c.Set("credential_id", credentialID)
return c.JSON(http.StatusOK, map[string]any{
"success": true,
"message": "Authentication successful",
"credentialId": credentialID,
"accessToken": authSession.AccessToken,
"expiresIn": 3600,
"userDID": userDID,
"sessionId": authSession.SessionID,
"authenticatorData": authenticatorData,
"signature": signature,
"userHandle": userHandle,
})
}
// generateChallenge creates a cryptographically secure challenge
func generateChallenge() (string, error) {
bytes := make([]byte, 32)
if _, err := rand.Read(bytes); err != nil {
return "", err
}
return base64.RawURLEncoding.EncodeToString(bytes), nil
}
// verifyClientData verifies client data JSON
func verifyClientData(clientDataJSON, expectedChallenge, expectedType string) error {
clientData, err := webauthn.ValidateClientDataJSONFormat(clientDataJSON)
if err != nil {
return fmt.Errorf("failed to parse client data: %w", err)
}
if clientData.Challenge != expectedChallenge {
return fmt.Errorf("challenge mismatch")
}
if clientData.Type != expectedType {
return fmt.Errorf(
"invalid client data type: expected %s, got %s",
expectedType,
clientData.Type,
)
}
// TODO: Verify origin from config
expectedOrigins := []string{
"http://localhost",
"http://localhost:8080",
"http://localhost:8081",
"http://localhost:8082",
"http://localhost:8083",
"http://localhost:8084",
}
validOrigin := false
for _, origin := range expectedOrigins {
if clientData.Origin == origin {
validOrigin = true
break
}
}
if !validOrigin {
return fmt.Errorf("invalid origin: %s", clientData.Origin)
}
return nil
}
// extractPublicKeyFromAttestation extracts public key from attestation object
func extractPublicKeyFromAttestation(attestationObjectB64 string) ([]byte, int32, error) {
// Validate format
if err := webauthn.ValidateAttestationObjectFormat(attestationObjectB64); err != nil {
return nil, 0, fmt.Errorf("invalid attestation format: %w", err)
}
// Decode attestation object
attestationBytes, err := base64.RawURLEncoding.DecodeString(attestationObjectB64)
if err != nil {
return nil, 0, fmt.Errorf("failed to decode attestation: %w", err)
}
// Parse CBOR
var attestationObj webauthn.AttestationObject
if err := webauthncbor.Unmarshal(attestationBytes, &attestationObj); err != nil {
return nil, 0, fmt.Errorf("failed to unmarshal attestation: %w", err)
}
// Unmarshal authenticator data
if err := attestationObj.AuthData.Unmarshal(attestationObj.RawAuthData); err != nil {
return nil, 0, fmt.Errorf("failed to unmarshal auth data: %w", err)
}
// Check for attested credential data
if !attestationObj.AuthData.Flags.HasAttestedCredentialData() {
return nil, 0, fmt.Errorf("no attested credential data")
}
publicKey := attestationObj.AuthData.AttData.CredentialPublicKey
if len(publicKey) == 0 {
return nil, 0, fmt.Errorf("no public key found")
}
// Default to ES256 algorithm
algorithm := int32(-7)
return publicKey, algorithm, nil
}
// GetWebAuthnCredentials retrieves credentials for a user
func GetWebAuthnCredentials(c echo.Context) error {
username := c.Param("username")
if username == "" {
return c.JSON(http.StatusBadRequest, map[string]string{
"error": "Username is required",
})
}
webAuthnStore.mu.RLock()
credentials, exists := webAuthnStore.credentials[username]
webAuthnStore.mu.RUnlock()
if !exists {
return c.JSON(http.StatusNotFound, map[string]string{
"error": "No credentials found for user",
})
}
// Return sanitized credentials (without sensitive data)
sanitized := make([]map[string]any, len(credentials))
for i, cred := range credentials {
sanitized[i] = map[string]any{
"credentialId": cred.CredentialID,
"createdAt": cred.CreatedAt,
"algorithm": cred.Algorithm,
}
}
return c.JSON(http.StatusOK, sanitized)
}
+284
View File
@@ -0,0 +1,284 @@
package handlers
import (
"fmt"
"log"
"net/http"
"sync"
"time"
"github.com/gorilla/websocket"
"github.com/labstack/echo/v4"
)
// TaskStatusMessage represents a WebSocket message for task status updates
type TaskStatusMessage struct {
TaskID string `json:"task_id"`
Status string `json:"status"`
Progress int `json:"progress,omitempty"`
Data any `json:"data,omitempty"`
Error string `json:"error,omitempty"`
Time time.Time `json:"timestamp"`
}
// ConnectionManager manages WebSocket connections for task status broadcasting
type ConnectionManager struct {
connections map[string]map[*websocket.Conn]bool // taskID -> connections
mutex sync.RWMutex
}
// NewConnectionManager creates a new connection manager
func NewConnectionManager() *ConnectionManager {
return &ConnectionManager{
connections: make(map[string]map[*websocket.Conn]bool),
mutex: sync.RWMutex{},
}
}
// AddConnection adds a WebSocket connection for a specific task ID
func (cm *ConnectionManager) AddConnection(taskID string, conn *websocket.Conn) {
cm.mutex.Lock()
defer cm.mutex.Unlock()
if cm.connections[taskID] == nil {
cm.connections[taskID] = make(map[*websocket.Conn]bool)
}
cm.connections[taskID][conn] = true
log.Printf("WebSocket connection added for task: %s", taskID)
}
// RemoveConnection removes a WebSocket connection
func (cm *ConnectionManager) RemoveConnection(taskID string, conn *websocket.Conn) {
cm.mutex.Lock()
defer cm.mutex.Unlock()
if connections, exists := cm.connections[taskID]; exists {
delete(connections, conn)
if len(connections) == 0 {
delete(cm.connections, taskID)
}
}
log.Printf("WebSocket connection removed for task: %s", taskID)
}
// BroadcastToTask broadcasts a message to all connections listening to a specific task
func (cm *ConnectionManager) BroadcastToTask(taskID string, message TaskStatusMessage) {
cm.mutex.RLock()
connections, exists := cm.connections[taskID]
cm.mutex.RUnlock()
if !exists {
return
}
for conn := range connections {
if err := conn.WriteJSON(message); err != nil {
log.Printf("Error broadcasting to WebSocket: %v", err)
cm.RemoveConnection(taskID, conn)
conn.Close()
}
}
}
// SSEManager manages Server-Sent Event connections for task status streaming
type SSEManager struct {
connections map[string]map[chan TaskStatusMessage]bool // taskID -> channels
mutex sync.RWMutex
}
// NewSSEManager creates a new SSE manager
func NewSSEManager() *SSEManager {
return &SSEManager{
connections: make(map[string]map[chan TaskStatusMessage]bool),
mutex: sync.RWMutex{},
}
}
// AddSSEConnection adds an SSE channel for a specific task ID
func (sm *SSEManager) AddSSEConnection(taskID string, ch chan TaskStatusMessage) {
sm.mutex.Lock()
defer sm.mutex.Unlock()
if sm.connections[taskID] == nil {
sm.connections[taskID] = make(map[chan TaskStatusMessage]bool)
}
sm.connections[taskID][ch] = true
log.Printf("SSE connection added for task: %s", taskID)
}
// RemoveSSEConnection removes an SSE channel
func (sm *SSEManager) RemoveSSEConnection(taskID string, ch chan TaskStatusMessage) {
sm.mutex.Lock()
defer sm.mutex.Unlock()
if channels, exists := sm.connections[taskID]; exists {
delete(channels, ch)
if len(channels) == 0 {
delete(sm.connections, taskID)
}
}
close(ch)
log.Printf("SSE connection removed for task: %s", taskID)
}
// BroadcastToSSE broadcasts a message to all SSE connections listening to a specific task
func (sm *SSEManager) BroadcastToSSE(taskID string, message TaskStatusMessage) {
sm.mutex.RLock()
channels, exists := sm.connections[taskID]
sm.mutex.RUnlock()
if !exists {
return
}
for ch := range channels {
select {
case ch <- message:
// Message sent successfully
default:
// Channel is blocked, remove it
log.Printf("SSE channel blocked, removing for task: %s", taskID)
sm.RemoveSSEConnection(taskID, ch)
}
}
}
// WebSocketHandler handles WebSocket connections for real-time task status updates
func WebSocketHandler(
upgrader *websocket.Upgrader,
connectionManager *ConnectionManager,
) echo.HandlerFunc {
return func(c echo.Context) error {
taskID := c.Param("task_id")
if taskID == "" {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Task ID is required"})
}
// Upgrade HTTP connection to WebSocket
ws, err := upgrader.Upgrade(c.Response(), c.Request(), nil)
if err != nil {
log.Printf("WebSocket upgrade failed: %v", err)
return err
}
defer ws.Close()
// Add connection to manager
connectionManager.AddConnection(taskID, ws)
defer connectionManager.RemoveConnection(taskID, ws)
// Send initial connection confirmation
initialMessage := TaskStatusMessage{
TaskID: taskID,
Status: "connected",
Time: time.Now(),
}
if err := ws.WriteJSON(initialMessage); err != nil {
log.Printf("Error sending initial message: %v", err)
return err
}
// Keep connection alive and handle incoming messages
for {
// Read message from client (ping/pong for keepalive)
_, _, err := ws.ReadMessage()
if err != nil {
log.Printf("WebSocket read error: %v", err)
break
}
// Echo back a pong message to keep connection alive
pongMessage := TaskStatusMessage{
TaskID: taskID,
Status: "pong",
Time: time.Now(),
}
if err := ws.WriteJSON(pongMessage); err != nil {
log.Printf("Error sending pong: %v", err)
break
}
}
return nil
}
}
// SSEHandler handles Server-Sent Events for streaming task status updates
func SSEHandler(sseManager *SSEManager) echo.HandlerFunc {
return func(c echo.Context) error {
taskID := c.Param("task_id")
if taskID == "" {
return c.JSON(http.StatusBadRequest, map[string]string{"error": "Task ID is required"})
}
// Set SSE headers
c.Response().Header().Set("Content-Type", "text/event-stream")
c.Response().Header().Set("Cache-Control", "no-cache")
c.Response().Header().Set("Connection", "keep-alive")
c.Response().Header().Set("Access-Control-Allow-Origin", "*")
c.Response().Header().Set("Access-Control-Allow-Headers", "Cache-Control")
// Create a channel for this SSE connection
messageCh := make(chan TaskStatusMessage, 10)
sseManager.AddSSEConnection(taskID, messageCh)
defer sseManager.RemoveSSEConnection(taskID, messageCh)
// Send initial connection message
initialMessage := TaskStatusMessage{
TaskID: taskID,
Status: "connected",
Time: time.Now(),
}
fmt.Fprintf(
c.Response().Writer,
"data: {\"task_id\":\"%s\",\"status\":\"connected\",\"timestamp\":\"%s\"}\n\n",
taskID,
initialMessage.Time.Format(time.RFC3339),
)
c.Response().Flush()
// Keep connection alive and send messages
for {
select {
case message, ok := <-messageCh:
if !ok {
return nil
}
// Format message as SSE data
sseData := fmt.Sprintf(
"data: {\"task_id\":\"%s\",\"status\":\"%s\",\"progress\":%d,\"timestamp\":\"%s\"}",
message.TaskID,
message.Status,
message.Progress,
message.Time.Format(time.RFC3339),
)
if message.Error != "" {
sseData = fmt.Sprintf(
"data: {\"task_id\":\"%s\",\"status\":\"%s\",\"error\":\"%s\",\"timestamp\":\"%s\"}",
message.TaskID,
message.Status,
message.Error,
message.Time.Format(time.RFC3339),
)
}
fmt.Fprintf(c.Response().Writer, "%s\n\n", sseData)
c.Response().Flush()
case <-c.Request().Context().Done():
// Client disconnected
return nil
case <-time.After(30 * time.Second):
// Send keepalive message every 30 seconds
fmt.Fprintf(
c.Response().Writer,
"data: {\"task_id\":\"%s\",\"status\":\"keepalive\",\"timestamp\":\"%s\"}\n\n",
taskID,
time.Now().Format(time.RFC3339),
)
c.Response().Flush()
}
}
}
}