diff --git a/internal/db/whitelist.go b/internal/db/whitelist.go index 7c2bb78..78fa0cd 100644 --- a/internal/db/whitelist.go +++ b/internal/db/whitelist.go @@ -8,7 +8,7 @@ type WhitelistEntry struct { ID uint `gorm:"primaryKey" json:"id"` DID string `gorm:"column:did;uniqueIndex;not null" json:"did"` Handle string `json:"handle"` - MaxNodes int `gorm:"default:0" json:"max_nodes"` // 0 = unlimited + Email string `json:"email"` Notes string `json:"notes"` CreatedAt time.Time `json:"created_at"` } diff --git a/oidc/provider.go b/oidc/provider.go index 462d07a..37822bc 100644 --- a/oidc/provider.go +++ b/oidc/provider.go @@ -65,6 +65,7 @@ func (p *Provider) Discovery() map[string]any { "name", "preferred_username", "email", + "email_verified", }, "code_challenge_methods_supported": []string{ "S256", @@ -110,6 +111,7 @@ func (p *Provider) IssueIDToken(sub string, preferredUsername string, email stri Expiration(now.Add(1 * time.Hour)). Claim("preferred_username", preferredUsername). Claim("email", email). + Claim("email_verified", true). Claim("name", preferredUsername). Claim("nonce", nonce). Build() diff --git a/server/admin.go b/server/admin.go index 57f031d..490e6d1 100644 --- a/server/admin.go +++ b/server/admin.go @@ -67,10 +67,10 @@ func (s *Server) handleListWhitelist(e echo.Context) error { // handleAddWhitelist adds a new whitelist entry. func (s *Server) handleAddWhitelist(e echo.Context) error { var input struct { - DID string `json:"did"` - Handle string `json:"handle"` - MaxNodes int `json:"max_nodes"` - Notes string `json:"notes"` + DID string `json:"did"` + Handle string `json:"handle"` + Email string `json:"email"` + Notes string `json:"notes"` } if err := e.Bind(&input); err != nil { return e.JSON(http.StatusBadRequest, map[string]string{"error": "invalid request"}) @@ -80,10 +80,10 @@ func (s *Server) handleAddWhitelist(e echo.Context) error { } entry := db.WhitelistEntry{ - DID: input.DID, - Handle: input.Handle, - MaxNodes: input.MaxNodes, - Notes: input.Notes, + DID: input.DID, + Handle: input.Handle, + Email: input.Email, + Notes: input.Notes, } if err := s.db.DB.Create(&entry).Error; err != nil { return e.JSON(http.StatusConflict, map[string]string{"error": "DID already exists"}) @@ -124,9 +124,9 @@ func (s *Server) handleUpdateWhitelist(e echo.Context) error { } var input struct { - Handle *string `json:"handle"` - MaxNodes *int `json:"max_nodes"` - Notes *string `json:"notes"` + Handle *string `json:"handle"` + Email *string `json:"email"` + Notes *string `json:"notes"` } if err := e.Bind(&input); err != nil { return e.JSON(http.StatusBadRequest, map[string]string{"error": "invalid request"}) @@ -135,8 +135,8 @@ func (s *Server) handleUpdateWhitelist(e echo.Context) error { if input.Handle != nil { entry.Handle = *input.Handle } - if input.MaxNodes != nil { - entry.MaxNodes = *input.MaxNodes + if input.Email != nil { + entry.Email = *input.Email } if input.Notes != nil { entry.Notes = *input.Notes diff --git a/server/callback.go b/server/callback.go index f178701..bc245aa 100644 --- a/server/callback.go +++ b/server/callback.go @@ -1,7 +1,6 @@ package server import ( - "encoding/json" "fmt" "net/http" "time" @@ -53,54 +52,19 @@ func (s *Server) handleATProtoCallback(e echo.Context) error { // Whitelist check: if the whitelist has any entries, only allow DIDs in the list. // Empty whitelist = allow all (bootstrap mode). + var entry db.WhitelistEntry var count int64 s.db.DB.Model(&db.WhitelistEntry{}).Count(&count) if count > 0 { - var entry db.WhitelistEntry if err := s.db.DB.Where("did = ?", did).First(&entry).Error; err != nil { s.logger.Info("access denied — DID not in whitelist", "did", did) return s.renderError(e, "Access Denied", "Your AT Protocol identity is not authorized to join this mesh.") } - s.logger.Info("access granted — DID in whitelist", "did", did, "handle", entry.Handle, "maxnodes", entry.MaxNodes) - - // Device limit enforcement: if MaxNodes > 0, check how many nodes - // the user already has registered in Headscale. - if entry.MaxNodes > 0 && s.headscale != nil { - nodesData, err := s.headscale.ListNodes() - if err != nil { - s.logger.Error("failed to list nodes for device limit check", "err", err) - return s.renderError(e, "Server Error", "Could not verify device limit.") - } - - var nodesResp struct { - Nodes []struct { - User struct { - Name string `json:"name"` - } `json:"user"` - } `json:"nodes"` - } - if err := json.Unmarshal(nodesData, &nodesResp); err != nil { - s.logger.Error("failed to parse nodes response", "err", err) - return s.renderError(e, "Server Error", "Could not verify device limit.") - } - - count := 0 - for _, node := range nodesResp.Nodes { - if node.User.Name == did { - count++ - } - } - - if count >= entry.MaxNodes { - s.logger.Info("device limit reached", "did", did, "current", count, "max", entry.MaxNodes) - return s.renderError(e, "Device Limit Reached", - fmt.Sprintf("You have reached the maximum number of devices (%d) for your account.", entry.MaxNodes)) - } - } + s.logger.Info("access granted — DID in whitelist", "did", did, "handle", entry.Handle, "email", entry.Email) } - // Issue an OIDC auth code for Headscale, using the real DID as sub + // Issue an OIDC auth code, using the real DID as sub oidcCode := oidc.GenerateAuthCode() authReq := &db.OidcAuthCode{ Code: oidcCode, @@ -113,7 +77,7 @@ func (s *Server) handleATProtoCallback(e echo.Context) error { CodeChallengeMethod: bridge.OidcCodeChallengeMethod, Sub: did, PreferredUsername: bridge.Handle, - Email: "", + Email: entry.Email, ExpiresAt: time.Now().Add(10 * time.Minute), } diff --git a/server/server_test.go b/server/server_test.go index ed58e54..13e6e6e 100644 --- a/server/server_test.go +++ b/server/server_test.go @@ -866,7 +866,7 @@ func TestWhitelistAdd(t *testing.T) { base := startTestServer(t, s) client := adminClient(t, base) - body := bytes.NewBufferString(`{"did":"did:plc:test123","handle":"test.bsky.social","max_nodes":3,"notes":"test entry"}`) + body := bytes.NewBufferString(`{"did":"did:plc:test123","handle":"test.bsky.social","email":"test@mesh.glados.computer","notes":"test entry"}`) req, _ := http.NewRequest("POST", base+"/api/v1/whitelist", body) req.Header.Set("Content-Type", "application/json") @@ -888,8 +888,8 @@ func TestWhitelistAdd(t *testing.T) { if entry.Handle != "test.bsky.social" { t.Errorf("handle = %v", entry.Handle) } - if entry.MaxNodes != 3 { - t.Errorf("max_nodes = %v", entry.MaxNodes) + if entry.Email != "test@mesh.glados.computer" { + t.Errorf("email = %v", entry.Email) } } @@ -900,10 +900,10 @@ func TestWhitelistList(t *testing.T) { // Add an entry first s.db.DB.Create(&db.WhitelistEntry{ - DID: "did:plc:listtest", - Handle: "list.bsky.social", - MaxNodes: 2, - Notes: "list test", + DID: "did:plc:listtest", + Handle: "list.bsky.social", + Email: "list@mesh.glados.computer", + Notes: "list test", }) resp, err := client.Get(base + "/api/v1/whitelist") @@ -932,9 +932,9 @@ func TestWhitelistDelete(t *testing.T) { client := adminClient(t, base) entry := db.WhitelistEntry{ - DID: "did:plc:deletetest", - Handle: "delete.bsky.social", - MaxNodes: 1, + DID: "did:plc:deletetest", + Handle: "delete.bsky.social", + Email: "delete@mesh.glados.computer", } s.db.DB.Create(&entry) @@ -964,14 +964,14 @@ func TestWhitelistUpdate(t *testing.T) { client := adminClient(t, base) entry := db.WhitelistEntry{ - DID: "did:plc:updatetest", - Handle: "old.bsky.social", - MaxNodes: 1, - Notes: "old notes", + DID: "did:plc:updatetest", + Handle: "old.bsky.social", + Email: "old@mesh.glados.computer", + Notes: "old notes", } s.db.DB.Create(&entry) - body := bytes.NewBufferString(`{"handle":"new.bsky.social","max_nodes":5,"notes":"updated notes"}`) + body := bytes.NewBufferString(`{"handle":"new.bsky.social","email":"new@mesh.glados.computer","notes":"updated notes"}`) req, _ := http.NewRequest("PUT", fmt.Sprintf("%s/api/v1/whitelist/%d", base, entry.ID), body) req.Header.Set("Content-Type", "application/json") @@ -990,8 +990,8 @@ func TestWhitelistUpdate(t *testing.T) { if updated.Handle != "new.bsky.social" { t.Errorf("handle = %v, want new.bsky.social", updated.Handle) } - if updated.MaxNodes != 5 { - t.Errorf("max_nodes = %v, want 5", updated.MaxNodes) + if updated.Email != "new@mesh.glados.computer" { + t.Errorf("email = %v, want new@mesh.glados.computer", updated.Email) } if updated.Notes != "updated notes" { t.Errorf("notes = %v, want 'updated notes'", updated.Notes) diff --git a/server/token.go b/server/token.go index 53d246f..d2a2a4d 100644 --- a/server/token.go +++ b/server/token.go @@ -128,9 +128,11 @@ func (s *Server) handleUserinfo(e echo.Context) error { return e.JSON(http.StatusUnauthorized, map[string]string{"error": "invalid_token"}) } - return e.JSON(http.StatusOK, map[string]string{ - "sub": authReq.Sub, - "name": authReq.PreferredUsername, + return e.JSON(http.StatusOK, map[string]any{ + "sub": authReq.Sub, + "name": authReq.PreferredUsername, "preferred_username": authReq.PreferredUsername, + "email": authReq.Email, + "email_verified": true, }) }