diff --git a/cmd/atlogin/main.go b/cmd/atlogin/main.go index 18215f8..be23c2f 100644 --- a/cmd/atlogin/main.go +++ b/cmd/atlogin/main.go @@ -585,27 +585,53 @@ func (s *idpServer) serveAuthorize(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, atprotoRedirectURL, http.StatusFound) } -// parseLoginHint extracts the ATProto handle and domain from a login hint in the format: -// handle@any.domain -> (handle, any.domain) -// We accept any domain and will service it as an authoritative login provider. +// parseLoginHint extracts the ATProto handle and domain from a login hint. +// +// Format: user@domain +// +// ATProto handle rules: +// 1. Default: user@domain -> @user.domain (ATProto handle), domain (webfinger domain) +// 2. If "user." is a prefix of domain: user@user.example.com -> @user.example.com, user.example.com +// 3. Special case for at.apenwarr.ca: user@at.apenwarr.ca -> @user, at.apenwarr.ca +// +// Examples: +// - at@apenwarr.ca -> @at.apenwarr.ca (handle), apenwarr.ca (domain) +// - apenwarr@apenwarr.ca -> @apenwarr.ca (handle), apenwarr.ca (domain) +// - user@at.apenwarr.ca -> @user (handle), at.apenwarr.ca (domain) [backward compat] +// +// Returns: (atprotoHandle, webfingerDomain, error) func parseLoginHint(loginHint string) (string, string, error) { parts := strings.SplitN(loginHint, "@", 2) if len(parts) != 2 { - return "", "", fmt.Errorf("expected format: handle@domain, got: %s", loginHint) + return "", "", fmt.Errorf("expected format: user@domain, got: %s", loginHint) } - handle := parts[0] + user := parts[0] domain := parts[1] - if handle == "" { - return "", "", fmt.Errorf("handle cannot be empty") + if user == "" { + return "", "", fmt.Errorf("user cannot be empty") } if domain == "" { return "", "", fmt.Errorf("domain cannot be empty") } - return handle, domain, nil + // Special case for backward compatibility with at.apenwarr.ca + if domain == "at.apenwarr.ca" { + return user, domain, nil + } + + // Check if "user." is a prefix of domain + prefix := user + "." + if strings.HasPrefix(domain, prefix) { + // user@user.example.com -> @user.example.com + return domain, domain, nil + } + + // Default case: user@domain -> @user.domain + atprotoHandle := user + "." + domain + return atprotoHandle, domain, nil } func (s *idpServer) serveToken(w http.ResponseWriter, r *http.Request) { diff --git a/cmd/atlogin/parse_test.go b/cmd/atlogin/parse_test.go new file mode 100644 index 0000000..4164baa --- /dev/null +++ b/cmd/atlogin/parse_test.go @@ -0,0 +1,84 @@ +package main + +import "testing" + +func TestParseLoginHint(t *testing.T) { + tests := []struct { + name string + loginHint string + wantHandle string + wantDomain string + wantErr bool + }{ + { + name: "at@apenwarr.ca -> @at.apenwarr.ca", + loginHint: "at@apenwarr.ca", + wantHandle: "at.apenwarr.ca", + wantDomain: "apenwarr.ca", + }, + { + name: "apenwarr@apenwarr.ca -> @apenwarr.ca", + loginHint: "apenwarr@apenwarr.ca", + wantHandle: "apenwarr.ca", + wantDomain: "apenwarr.ca", + }, + { + name: "user@at.apenwarr.ca -> @user (backward compat)", + loginHint: "user@at.apenwarr.ca", + wantHandle: "user", + wantDomain: "at.apenwarr.ca", + }, + { + name: "alice@example.com -> @alice.example.com", + loginHint: "alice@example.com", + wantHandle: "alice.example.com", + wantDomain: "example.com", + }, + { + name: "john@john.doe.com -> @john.doe.com (prefix match)", + loginHint: "john@john.doe.com", + wantHandle: "john.doe.com", + wantDomain: "john.doe.com", + }, + { + name: "hello@hello.example.com -> @hello.example.com (prefix match)", + loginHint: "hello@hello.example.com", + wantHandle: "hello.example.com", + wantDomain: "hello.example.com", + }, + { + name: "empty user", + loginHint: "@example.com", + wantErr: true, + }, + { + name: "no @ sign", + loginHint: "invalid", + wantErr: true, + }, + { + name: "empty domain", + loginHint: "user@", + wantErr: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gotHandle, gotDomain, err := parseLoginHint(tt.loginHint) + if (err != nil) != tt.wantErr { + t.Errorf("parseLoginHint() error = %v, wantErr %v", err, tt.wantErr) + return + } + if err != nil { + return + } + if gotHandle != tt.wantHandle { + t.Errorf("parseLoginHint() handle = %v, want %v", gotHandle, tt.wantHandle) + } + if gotDomain != tt.wantDomain { + t.Errorf("parseLoginHint() domain = %v, want %v", gotDomain, tt.wantDomain) + } + }) + } +}