From 0ad3d7fdb9f69bc22b2ef6a839cad55864d73ea9 Mon Sep 17 00:00:00 2001 From: Akshay Date: Tue, 4 Feb 2025 11:42:38 +0000 Subject: [PATCH] refactor signed client --- appview/state/signer.go | 102 ++++++++++++++++++++++++++++++++++++---- appview/state/state.go | 73 ++++++++++++++++------------ knotserver/handler.go | 2 +- knotserver/jetstream.go | 11 +++-- rbac/rbac.go | 31 ++++++++++++ 5 files changed, 176 insertions(+), 43 deletions(-) diff --git a/appview/state/signer.go b/appview/state/signer.go index 8a24f175..dc629fc3 100644 --- a/appview/state/signer.go +++ b/appview/state/signer.go @@ -1,10 +1,14 @@ package state import ( + "bytes" "crypto/hmac" "crypto/sha256" "encoding/hex" + "encoding/json" + "fmt" "net/http" + "net/url" "time" ) @@ -12,15 +16,6 @@ type SignerTransport struct { Secret string } -func SignedClient(secret string) *http.Client { - return &http.Client{ - Timeout: 5 * time.Second, - Transport: SignerTransport{ - Secret: secret, - }, - } -} - func (s SignerTransport) RoundTrip(req *http.Request) (*http.Response, error) { timestamp := time.Now().Format(time.RFC3339) mac := hmac.New(sha256.New, []byte(s.Secret)) @@ -31,3 +26,92 @@ func (s SignerTransport) RoundTrip(req *http.Request) (*http.Response, error) { req.Header.Set("X-Timestamp", timestamp) return http.DefaultTransport.RoundTrip(req) } + +type SignedClient struct { + Secret string + Url *url.URL + client *http.Client +} + +func NewSignedClient(domain, secret string) (*SignedClient, error) { + client := &http.Client{ + Timeout: 5 * time.Second, + Transport: SignerTransport{ + Secret: secret, + }, + } + + url, err := url.Parse(fmt.Sprintf("http://%s", domain)) + if err != nil { + return nil, err + } + + signedClient := &SignedClient{ + Secret: secret, + client: client, + Url: url, + } + + return signedClient, nil +} + +func (s *SignedClient) newRequest(method, endpoint string, body []byte) (*http.Request, error) { + return http.NewRequest(method, s.Url.JoinPath(endpoint).String(), bytes.NewReader(body)) +} + +func (s *SignedClient) Init(did string, keys []string) (*http.Response, error) { + const ( + Method = "POST" + Endpoint = "/init" + ) + + body, _ := json.Marshal(map[string]interface{}{ + "did": did, + "keys": keys, + }) + + req, err := s.newRequest(Method, Endpoint, body) + if err != nil { + return nil, err + } + + return s.client.Do(req) +} + +func (s *SignedClient) NewRepo(did, repoName string) (*http.Response, error) { + const ( + Method = "PUT" + Endpoint = "/repo/new" + ) + + body, _ := json.Marshal(map[string]interface{}{ + "did": did, + "name": repoName, + }) + + req, err := s.newRequest(Method, Endpoint, body) + if err != nil { + return nil, err + } + + return s.client.Do(req) +} + +func (s *SignedClient) AddMember(did string, keys []string) (*http.Response, error) { + const ( + Method = "PUT" + Endpoint = "/member/add" + ) + + body, _ := json.Marshal(map[string]interface{}{ + "did": did, + "keys": keys, + }) + + req, err := s.newRequest(Method, Endpoint, body) + if err != nil { + return nil, err + } + + return s.client.Do(req) +} diff --git a/appview/state/state.go b/appview/state/state.go index 27b256a9..b0ab59ba 100644 --- a/appview/state/state.go +++ b/appview/state/state.go @@ -1,14 +1,13 @@ package state import ( - "bytes" "crypto/hmac" "crypto/sha256" "encoding/hex" - "encoding/json" "fmt" "log" "net/http" + "path/filepath" "strings" "time" @@ -234,26 +233,18 @@ func (s *State) InitKnotServer(w http.ResponseWriter, r *http.Request) { } log.Println("checking ", domain) - url := fmt.Sprintf("http://%s/init", domain) - - body, _ := json.Marshal(map[string]interface{}{ - "did": user.Did, - "keys": []string{}, - }) - pingRequest, err := http.NewRequest("POST", url, bytes.NewBuffer(body)) + secret, err := s.db.GetRegistrationKey(domain) if err != nil { - log.Println("failed to build ping request", err) + log.Printf("no key found for domain %s: %s\n", domain, err) return } - secret, err := s.db.GetRegistrationKey(domain) + client, err := NewSignedClient(domain, secret) if err != nil { - log.Printf("no key found for domain %s: %s\n", domain, err) - return + log.Println("failed to create client to ", domain) } - client := SignedClient(secret) - resp, err := client.Do(pingRequest) + resp, err := client.Init(user.Did, []string{}) if err != nil { w.Write([]byte("no dice")) log.Println("domain was unreachable after 5 seconds") @@ -338,14 +329,14 @@ func (s *State) KnotServerInfo(w http.ResponseWriter, r *http.Request) { var members []string if reg.Registered != nil { - members, err = s.enforcer.E.GetUsersForRole("server:member", domain) + members, err = s.enforcer.GetUserByRole("server:member", domain) if err != nil { w.Write([]byte("failed to fetch member list")) return } } - ok, err := s.enforcer.E.HasGroupingPolicy(user.Did, "server:owner", domain) + ok, err := s.enforcer.IsServerOwner(user.Did, domain) isOwner := err == nil && ok p := pages.KnotParams{ @@ -382,7 +373,7 @@ func (s *State) ListMembers(w http.ResponseWriter, r *http.Request) { } // list all members for this domain - memberDids, err := s.enforcer.E.GetUsersForRole("server:member", domain) + memberDids, err := s.enforcer.GetUserByRole("server:member", domain) if err != nil { w.Write([]byte("failed to fetch member list")) return @@ -433,9 +424,30 @@ func (s *State) AddMember(w http.ResponseWriter, r *http.Request) { log.Printf("failed to create record: %s", err) return } - log.Println("created atproto record: ", resp.Uri) + secret, err := s.db.GetRegistrationKey(domain) + if err != nil { + log.Printf("no key found for domain %s: %s\n", domain, err) + return + } + + ksClient, err := NewSignedClient(domain, secret) + if err != nil { + log.Println("failed to create client to ", domain) + return + } + + ksResp, err := ksClient.AddMember(memberIdent.DID.String(), []string{}) + if err != nil { + log.Printf("failet to make request to %s: %s", domain, err) + } + + if ksResp.StatusCode != http.StatusNoContent { + w.Write([]byte(fmt.Sprint("knotserver failed to add member: ", err))) + return + } + err = s.enforcer.AddMember(domain, memberIdent.DID.String()) if err != nil { w.Write([]byte(fmt.Sprint("failed to add member: ", err))) @@ -481,21 +493,16 @@ func (s *State) AddRepo(w http.ResponseWriter, r *http.Request) { return } - client := SignedClient(secret) - url := fmt.Sprintf("http://%s/repo/new", domain) - body, _ := json.Marshal(map[string]interface{}{ - "did": user.Did, - "name": repoName, - }) - createRepoRequest, err := http.NewRequest("PUT", url, bytes.NewReader(body)) - - resp, err := client.Do(createRepoRequest) + client, err := NewSignedClient(domain, secret) + if err != nil { + log.Println("failed to create client to ", domain) + } + resp, err := client.NewRepo(user.Did, repoName) if err != nil { log.Println("failed to send create repo request", err) return } - if resp.StatusCode != http.StatusNoContent { log.Println("server returned ", resp.StatusCode) return @@ -507,13 +514,19 @@ func (s *State) AddRepo(w http.ResponseWriter, r *http.Request) { Name: repoName, Knot: domain, } - err = s.db.AddRepo(repo) if err != nil { log.Println("failed to add repo to db", err) return } + // acls + err = s.enforcer.AddRepo(user.Did, domain, filepath.Join(user.Did, repoName)) + if err != nil { + log.Println("failed to set up acls", err) + return + } + w.Write([]byte("created!")) } } diff --git a/knotserver/handler.go b/knotserver/handler.go index bc02182f..e99074bb 100644 --- a/knotserver/handler.go +++ b/knotserver/handler.go @@ -92,7 +92,7 @@ func Setup(ctx context.Context, c *config.Config, db *db.DB, e *rbac.Enforcer) ( r.Route("/member", func(r chi.Router) { r.Use(h.VerifySignature) - r.Put("/add", h.NewRepo) + r.Put("/add", h.AddMember) }) // Initialize the knot with an owner and public key. diff --git a/knotserver/jetstream.go b/knotserver/jetstream.go index b16b4b67..aec66dc8 100644 --- a/knotserver/jetstream.go +++ b/knotserver/jetstream.go @@ -7,7 +7,7 @@ import ( "io" "log" "net/http" - "path" + "net/url" "strings" "time" @@ -66,7 +66,13 @@ func (h *Handle) processPublicKey(did string, record map[string]interface{}) { } func (h *Handle) fetchAndAddKeys(did string) { - resp, err := http.Get(path.Join(h.c.AppViewEndpoint, did)) + keysEndpoint, err := url.JoinPath(h.c.AppViewEndpoint, "keys", did) + if err != nil { + log.Printf("error building endpoint url: %s: %v", did, err) + return + } + + resp, err := http.Get(keysEndpoint) if err != nil { log.Printf("error getting keys for %s: %v", did, err) return @@ -112,7 +118,6 @@ func (h *Handle) processKnotMember(did string, record map[string]interface{}) { } func (h *Handle) processMessages(messages <-chan []byte) { - log.Println("waiting for knot to be initialized") <-h.init log.Println("initalized jetstream watcher") diff --git a/rbac/rbac.go b/rbac/rbac.go index 32546243..6bf781e7 100644 --- a/rbac/rbac.go +++ b/rbac/rbac.go @@ -3,6 +3,7 @@ package rbac import ( "database/sql" "path" + "strings" sqladapter "github.com/Blank-Xu/sql-adapter" "github.com/casbin/casbin/v2" @@ -100,6 +101,36 @@ func (e *Enforcer) AddRepo(member, domain, repo string) error { return err } +func (e *Enforcer) GetUserByRole(role, domain string) ([]string, error) { + var membersWithoutRoles []string + + // this includes roles too, casbin does not differentiate. + // the filtering criteria is to remove strings not starting with `did:` + members, err := e.E.Enforcer.GetImplicitUsersForRole(role, domain) + for _, m := range members { + if strings.HasPrefix(m, "did:") { + membersWithoutRoles = append(membersWithoutRoles, m) + } + } + if err != nil { + return nil, err + } + + return membersWithoutRoles, nil +} + +func (e *Enforcer) isRole(user, role, domain string) (bool, error) { + return e.E.HasGroupingPolicy(user, role, domain) +} + +func (e *Enforcer) IsServerOwner(user, domain string) (bool, error) { + return e.isRole(user, "server:owner", domain) +} + +func (e *Enforcer) IsServerMember(user, domain string) (bool, error) { + return e.isRole(user, "server:member", domain) +} + // keyMatch2Func is a wrapper for keyMatch2 to make it compatible with Casbin func keyMatch2Func(args ...interface{}) (interface{}, error) { name1 := args[0].(string) -- 2.51.2