From e1fafe7c93e3f3aa378aae404804dbbab3cc4dd6 Mon Sep 17 00:00:00 2001 From: klizhentas Date: Tue, 30 Jun 2015 16:12:18 -0700 Subject: [PATCH] WIP: adding backend sessions support --- auth/api.go | 5 +- auth/clt.go | 59 ++++++++++++++++++++-- auth/srv.go | 92 +++++++++++++++++++++++++++++---- auth/srv_test.go | 80 ++++++++++++++++++++++------- auth/tun_test.go | 5 +- backend/backend.go | 11 ++++ backend/boltbk/boltbk.go | 91 ++++++++++++++++++++++++++++++++- backend/boltbk/boltbk_test.go | 8 +++ backend/codec.go | 30 +++++++++++ backend/etcdbk/etcd.go | 34 +++++++++++++ backend/etcdbk/etcd_test.go | 8 +++ backend/membk/mem.go | 20 ++++++++ backend/test/suite.go | 36 +++++++++++++ session/session.go | 89 ++++++++++++++++++++++++++++++++ session/session_test.go | 95 +++++++++++++++++++++++++++++++++++ 15 files changed, 627 insertions(+), 36 deletions(-) create mode 100644 backend/codec.go create mode 100644 session/session.go create mode 100644 session/session_test.go diff --git a/auth/api.go b/auth/api.go index 12f8c933e57..08d39f4fba3 100644 --- a/auth/api.go +++ b/auth/api.go @@ -7,16 +7,17 @@ import ( "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/memlog" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/mailgun/log" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/mailgun/oxy/trace" + "github.com/gravitational/teleport/session" "github.com/gravitational/teleport/utils" ) -func StartHTTPServer(a string, srv *AuthServer) error { +func StartHTTPServer(a string, srv *AuthServer, se session.SessionServer) error { addr, err := utils.ParseAddr(a) if err != nil { return err } t, err := trace.New( - NewAPIServer(srv, memlog.New()), + NewAPIServer(srv, memlog.New(), se), log.GetLogger().Writer(log.SeverityInfo)) if err != nil { return err diff --git a/auth/clt.go b/auth/clt.go index a37f8e1fa61..f662061bf57 100644 --- a/auth/clt.go +++ b/auth/clt.go @@ -10,6 +10,7 @@ import ( "time" "github.com/gravitational/teleport/backend" + "github.com/gravitational/teleport/session" "github.com/gravitational/teleport/utils" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/roundtrip" @@ -85,6 +86,58 @@ func (c *Client) Delete(u string) (*roundtrip.Response, error) { return c.convertResponse(c.Client.Delete(u)) } +func (c *Client) GetSessions() ([]session.Session, error) { + out, err := c.Get(c.Endpoint("sessions"), url.Values{}) + if err != nil { + return nil, err + } + var re *sessionsResponse + if err := json.Unmarshal(out.Bytes(), &re); err != nil { + return nil, err + } + return re.Sessions, nil +} + +func (c *Client) GetSession(id string) (*session.Session, error) { + out, err := c.Get(c.Endpoint("sessions", id), url.Values{}) + if err != nil { + return nil, err + } + var re *sessionResponse + if err := json.Unmarshal(out.Bytes(), &re); err != nil { + return nil, err + } + return &re.Session, nil +} + +func (c *Client) DeleteSession(id string) error { + _, err := c.Delete(c.Endpoint("sessions", id)) + return err +} + +func (c *Client) UpsertParty(id string, p session.Party, ttl time.Duration) error { + a, err := p.LastActive.MarshalText() + if err != nil { + return err + } + out, err := c.PostForm(c.Endpoint("sessions", id, "parties"), url.Values{ + "id": []string{p.ID}, + "site": []string{p.Site}, + "user": []string{p.User}, + "server": []string{p.Server}, + "ttl": []string{ttl.String()}, + "last_active": []string{string(a)}, + }) + if err != nil { + return err + } + var re *partyResponse + if err := json.Unmarshal(out.Bytes(), &re); err != nil { + return err + } + return nil +} + func (c *Client) UpsertRemoteCert(cert backend.RemoteCert, ttl time.Duration) error { out, err := c.PostForm(c.Endpoint("ca", "remote", cert.Type, "hosts", cert.FQDN), url.Values{ "key": []string{string(cert.Value)}, @@ -267,7 +320,7 @@ func (c *Client) SignIn(user string, password []byte) (string, error) { if err != nil { return "", err } - var re *sessionResponse + var re *webSessionResponse if err := json.Unmarshal(out.Bytes(), &re); err != nil { return "", err } @@ -282,7 +335,7 @@ func (c *Client) GetWebSession(user string, sid string) (string, error) { if err != nil { return "", err } - var re *sessionResponse + var re *webSessionResponse if err := json.Unmarshal(out.Bytes(), &re); err != nil { return "", err } @@ -300,7 +353,7 @@ func (c *Client) GetWebSessionsKeys( if err != nil { return nil, err } - var re *sessionsResponse + var re *webSessionsResponse if err := json.Unmarshal(out.Bytes(), &re); err != nil { return nil, err } diff --git a/auth/srv.go b/auth/srv.go index 5cecfdbb2f6..4e47c743a5d 100644 --- a/auth/srv.go +++ b/auth/srv.go @@ -10,10 +10,11 @@ import ( "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/form" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/memlog" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/roundtrip" - "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/session" + websession "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/session" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/julienschmidt/httprouter" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/mailgun/log" "github.com/gravitational/teleport/backend" + "github.com/gravitational/teleport/session" ) type Config struct { @@ -26,12 +27,14 @@ type APIServer struct { httprouter.Router s *AuthServer elog memlog.Logger + se session.SessionServer } -func NewAPIServer(s *AuthServer, elog memlog.Logger) *APIServer { +func NewAPIServer(s *AuthServer, elog memlog.Logger, se session.SessionServer) *APIServer { srv := &APIServer{ s: s, elog: elog, + se: se, } srv.Router = *httprouter.New() @@ -89,6 +92,12 @@ func NewAPIServer(s *AuthServer, elog memlog.Logger) *APIServer { srv.POST("/v1/events", srv.submitEvents) srv.GET("/v1/events", srv.getEvents) + // Sesssions + srv.POST("/v1/sessions/:id/parties", srv.upsertSessionParty) + srv.GET("/v1/sessions", srv.getSessions) + srv.GET("/v1/sessions/:id", srv.getSession) + srv.DELETE("/v1/sessions/:id", srv.deleteSession) + return srv } @@ -178,7 +187,7 @@ func (s *APIServer) getWebTuns(w http.ResponseWriter, r *http.Request, p httprou func (s *APIServer) deleteWebSession(w http.ResponseWriter, r *http.Request, p httprouter.Params) { user, sid := p[0].Value, p[1].Value - err := s.s.DeleteWebSession(user, session.SecureID(sid)) + err := s.s.DeleteWebSession(user, websession.SecureID(sid)) if err != nil { replyErr(w, err) return @@ -188,12 +197,12 @@ func (s *APIServer) deleteWebSession(w http.ResponseWriter, r *http.Request, p h func (s *APIServer) getWebSession(w http.ResponseWriter, r *http.Request, p httprouter.Params) { user, sid := p[0].Value, p[1].Value - ws, err := s.s.GetWebSession(user, session.SecureID(sid)) + ws, err := s.s.GetWebSession(user, websession.SecureID(sid)) if err != nil { replyErr(w, err) return } - reply(w, http.StatusOK, &sessionResponse{SID: string(ws.SID)}) + reply(w, http.StatusOK, &webSessionResponse{SID: string(ws.SID)}) } func (s *APIServer) getWebSessions(w http.ResponseWriter, r *http.Request, p httprouter.Params) { @@ -203,7 +212,7 @@ func (s *APIServer) getWebSessions(w http.ResponseWriter, r *http.Request, p htt replyErr(w, err) return } - reply(w, http.StatusOK, &sessionsResponse{Keys: keys}) + reply(w, http.StatusOK, &webSessionsResponse{Keys: keys}) } func (s *APIServer) signIn(w http.ResponseWriter, r *http.Request, p httprouter.Params) { @@ -221,7 +230,7 @@ func (s *APIServer) signIn(w http.ResponseWriter, r *http.Request, p httprouter. replyErr(w, err) return } - reply(w, http.StatusOK, &sessionResponse{SID: string(ws.SID)}) + reply(w, http.StatusOK, &webSessionResponse{SID: string(ws.SID)}) } func (s *APIServer) upsertPassword(w http.ResponseWriter, r *http.Request, p httprouter.Params) { @@ -564,6 +573,59 @@ func (s *APIServer) deleteRemoteCert(w http.ResponseWriter, r *http.Request, p h reply(w, http.StatusOK, message(fmt.Sprintf("cert '%v' deleted", id))) } +func (s *APIServer) upsertSessionParty(w http.ResponseWriter, r *http.Request, p httprouter.Params) { + var party session.Party + + sid := p[0].Value + var ttl time.Duration + + err := form.Parse(r, + form.String("id", &party.ID, form.Required()), + form.String("site", &party.Site, form.Required()), + form.String("user", &party.User, form.Required()), + form.String("server", &party.Server, form.Required()), + form.Duration("ttl", &ttl), + ) + if err != nil { + replyErr(w, err) + return + } + party.LastActive = time.Now() + if err := s.se.UpsertParty(sid, party, ttl); err != nil { + replyErr(w, err) + return + } + reply(w, http.StatusOK, partyResponse{Party: party}) +} + +func (s *APIServer) getSessions(w http.ResponseWriter, r *http.Request, p httprouter.Params) { + sessions, err := s.se.GetSessions() + if err != nil { + replyErr(w, err) + return + } + reply(w, http.StatusOK, &sessionsResponse{Sessions: sessions}) +} + +func (s *APIServer) getSession(w http.ResponseWriter, r *http.Request, p httprouter.Params) { + sid := p[0].Value + se, err := s.se.GetSession(sid) + if err != nil { + replyErr(w, err) + return + } + reply(w, http.StatusOK, &sessionResponse{Session: *se}) +} + +func (s *APIServer) deleteSession(w http.ResponseWriter, r *http.Request, p httprouter.Params) { + sid := p[0].Value + if err := s.se.DeleteSession(sid); err != nil { + replyErr(w, err) + return + } + reply(w, http.StatusOK, message(fmt.Sprintf("session %v was deleted", sid))) +} + type pubKeyResponse struct { PubKey string `json:"pubkey"` } @@ -593,11 +655,11 @@ type keyPairResponse struct { PubKey string `json:"pubkey"` } -type sessionResponse struct { +type webSessionResponse struct { SID string `json:"sid"` } -type sessionsResponse struct { +type webSessionsResponse struct { Keys []backend.AuthorizedKey `json:"keys"` } @@ -621,6 +683,18 @@ type eventsResponse struct { Events []interface{} `json:"events"` } +type partyResponse struct { + Party session.Party `json:"party"` +} + +type sessionsResponse struct { + Sessions []session.Session `json:"sessions"` +} + +type sessionResponse struct { + Session session.Session `json:"session"` +} + func message(msg string) map[string]interface{} { return map[string]interface{}{"message": msg} } diff --git a/auth/srv_test.go b/auth/srv_test.go index 129e44412d0..213d70d18b5 100644 --- a/auth/srv_test.go +++ b/auth/srv_test.go @@ -2,13 +2,15 @@ package auth import ( "net/http/httptest" + "path/filepath" "testing" "time" "github.com/gravitational/teleport/Godeps/_workspace/src/golang.org/x/crypto/ssh" authority "github.com/gravitational/teleport/auth/native" "github.com/gravitational/teleport/backend" - "github.com/gravitational/teleport/backend/membk" + "github.com/gravitational/teleport/backend/boltbk" + "github.com/gravitational/teleport/session" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/gravitational/memlog" "github.com/gravitational/teleport/Godeps/_workspace/src/github.com/mailgun/lemma/secret" @@ -20,9 +22,10 @@ func TestAPI(t *testing.T) { TestingT(t) } type APISuite struct { srv *httptest.Server clt *Client - bk *membk.MemBackend + bk *boltbk.BoltBackend scrt *secret.Service a *AuthServer + dir string } var _ = Suite(&APISuite{}) @@ -36,9 +39,13 @@ func (s *APISuite) SetUpSuite(c *C) { } func (s *APISuite) SetUpTest(c *C) { - s.bk = membk.New() + s.dir = c.MkDir() + var err error + s.bk, err = boltbk.New(filepath.Join(s.dir, "db")) + c.Assert(err, IsNil) + s.a = NewAuthServer(s.bk, authority.New(), s.scrt) - s.srv = httptest.NewServer(NewAPIServer(s.a, memlog.New())) + s.srv = httptest.NewServer(NewAPIServer(s.a, memlog.New(), session.New(s.bk))) clt, err := NewClient(s.srv.URL) c.Assert(err, IsNil) s.clt = clt @@ -46,28 +53,42 @@ func (s *APISuite) SetUpTest(c *C) { func (s *APISuite) TearDownTest(c *C) { s.srv.Close() + s.bk.Close() } func (s *APISuite) TestHostCACRUD(c *C) { c.Assert(s.clt.ResetHostCA(), IsNil) - ca := s.bk.HostCA + + hca, err := s.bk.GetHostCA() + c.Assert(err, IsNil) + c.Assert(s.clt.ResetHostCA(), IsNil) - c.Assert(ca, Not(DeepEquals), s.bk.HostCA) + + hca2, err := s.bk.GetHostCA() + c.Assert(err, IsNil) + + c.Assert(hca, Not(DeepEquals), hca2) key, err := s.clt.GetHostCAPub() c.Assert(err, IsNil) - c.Assert(key, DeepEquals, s.bk.HostCA.Pub) + c.Assert(key, DeepEquals, hca2.Pub) } func (s *APISuite) TestUserCACRUD(c *C) { c.Assert(s.clt.ResetUserCA(), IsNil) - ca := s.bk.UserCA + + uca, err := s.bk.GetUserCA() + c.Assert(err, IsNil) + c.Assert(s.clt.ResetUserCA(), IsNil) - c.Assert(ca, Not(DeepEquals), s.bk.UserCA) + uca2, err := s.bk.GetUserCA() + c.Assert(err, IsNil) + + c.Assert(uca, Not(DeepEquals), uca2) key, err := s.clt.GetUserCAPub() c.Assert(err, IsNil) - c.Assert(key, DeepEquals, s.bk.UserCA.Pub) + c.Assert(key, DeepEquals, uca2.Pub) } func (s *APISuite) TestGenerateKeyPair(c *C) { @@ -133,19 +154,18 @@ func (s *APISuite) TestUserKeyCRUD(c *C) { key := backend.AuthorizedKey{ID: "id", Value: pub} cert, err := s.clt.UpsertUserKey("user1", key, 0) c.Assert(err, IsNil) - c.Assert(string(s.bk.Users["user1"].Keys["id"].Value), DeepEquals, string(cert)) + + keys, err := s.bk.GetUserKeys("user1") + c.Assert(err, IsNil) + c.Assert(string(keys[0].Value), DeepEquals, string(cert)) _, _, _, _, err = ssh.ParseAuthorizedKey(cert) c.Assert(err, IsNil) - keys, err := s.clt.GetUserKeys("user1") - c.Assert(err, IsNil) - c.Assert(len(keys), Equals, 1) - c.Assert(string(keys[0].Value), DeepEquals, string(cert)) - c.Assert(s.clt.DeleteUserKey("user1", "id"), IsNil) - _, ok := s.bk.Users["user1"].Keys["id"] - c.Assert(ok, Equals, false) + keys, err = s.bk.GetUserKeys("user1") + c.Assert(err, IsNil) + c.Assert(len(keys), Equals, 0) } func (s *APISuite) TestPasswordCRUD(c *C) { @@ -256,7 +276,6 @@ func (s *APISuite) TestTokens(c *C) { } func (s *APISuite) TestRemoteCACRUD(c *C) { - key := backend.RemoteCert{ FQDN: "example.com", ID: "id", @@ -276,3 +295,26 @@ func (s *APISuite) TestRemoteCACRUD(c *C) { err = s.clt.DeleteRemoteCert(key.Type, key.FQDN, key.ID) c.Assert(err, FitsTypeOf, &backend.NotFoundError{}) } + +func (s *APISuite) TestSharedSessions(c *C) { + out, err := s.clt.GetSessions() + c.Assert(err, IsNil) + c.Assert(out, DeepEquals, []session.Session{}) + + p1 := session.Party{ + ID: "p1", + User: "bob", + Site: "example.com", + Server: "localhost:1", + LastActive: time.Date(2009, time.November, 10, 23, 0, 0, 0, time.UTC), + } + c.Assert(s.clt.UpsertParty("s1", p1, 0), IsNil) + + out, err = s.clt.GetSessions() + c.Assert(err, IsNil) + sess := session.Session{ + ID: "s1", + Parties: []session.Party{p1}, + } + c.Assert(out, DeepEquals, []session.Session{sess}) +} diff --git a/auth/tun_test.go b/auth/tun_test.go index d6ec1275f87..647a5856a98 100644 --- a/auth/tun_test.go +++ b/auth/tun_test.go @@ -11,6 +11,7 @@ import ( authority "github.com/gravitational/teleport/auth/native" "github.com/gravitational/teleport/backend" "github.com/gravitational/teleport/backend/membk" + "github.com/gravitational/teleport/session" "github.com/gravitational/teleport/sshutils" "github.com/gravitational/teleport/utils" @@ -49,7 +50,7 @@ func (s *TunSuite) TearDownTest(c *C) { func (s *TunSuite) SetUpTest(c *C) { s.bk = membk.New() s.a = NewAuthServer(s.bk, authority.New(), s.scrt) - s.srv = httptest.NewServer(NewAPIServer(s.a, memlog.New())) + s.srv = httptest.NewServer(NewAPIServer(s.a, memlog.New(), session.New(s.bk))) // set up host private key and certificate c.Assert(s.a.ResetHostCA(""), IsNil) @@ -82,7 +83,7 @@ func (s *TunSuite) TestUnixServerClient(c *C) { l, err := net.Listen("unix", socketPath) c.Assert(err, IsNil) - h := NewAPIServer(s.a, memlog.New()) + h := NewAPIServer(s.a, memlog.New(), session.New(s.bk)) srv := &httptest.Server{ Listener: l, Config: &http.Server{ diff --git a/backend/backend.go b/backend/backend.go index 71c17cbcf91..64ce51ef120 100644 --- a/backend/backend.go +++ b/backend/backend.go @@ -10,6 +10,12 @@ import ( // TODO(klizhentas) this is bloated. Split it into little backend interfaces // Backend represents configuration backend implementation for Teleport type Backend interface { + GetKeys(path []string) ([]string, error) + UpsertVal(path []string, key string, val []byte, ttl time.Duration) error + GetVal(path []string, key string) ([]byte, error) + DeleteKey(path []string, key string) error + DeleteBucket(path []string, bkt string) error + // Grab a lock that will be released automatically in ttl time AcquireLock(token string, ttl time.Duration) error @@ -202,3 +208,8 @@ const ( HostCert = "host" UserCert = "user" ) + +func IsNotFound(e error) bool { + _, ok := e.(*NotFoundError) + return ok +} diff --git a/backend/boltbk/boltbk.go b/backend/boltbk/boltbk.go index 0fc835b97f9..bea54381a5f 100644 --- a/backend/boltbk/boltbk.go +++ b/backend/boltbk/boltbk.go @@ -4,6 +4,7 @@ package boltbk import ( "encoding/json" "fmt" + "sort" "sync" "time" @@ -30,6 +31,69 @@ func New(path string) (*BoltBackend, error) { }, nil } +func (b *BoltBackend) GetKeys(path []string) ([]string, error) { + keys, err := b.getKeys(path) + if err != nil { + if isNotFound(err) { + return []string{}, nil + } + return nil, err + } + // now do an iteration to expire keys + for _, key := range keys { + b.GetVal(path, key) + } + keys, err = b.getKeys(path) + if err != nil { + if isNotFound(err) { + return []string{}, nil + } + return nil, err + } + sort.Sort(sort.StringSlice(keys)) + return keys, nil +} + +func (b *BoltBackend) UpsertVal(path []string, key string, val []byte, ttl time.Duration) error { + v := &kv{ + Created: time.Now(), + Value: val, + TTL: ttl, + } + bytes, err := json.Marshal(v) + if err != nil { + return err + } + return b.upsertKey(path, key, bytes) +} + +func (b *BoltBackend) GetVal(path []string, key string) ([]byte, error) { + var val []byte + if err := b.getKey(path, key, &val); err != nil { + return nil, err + } + var k *kv + if err := json.Unmarshal(val, &k); err != nil { + return nil, err + } + if k.TTL != 0 && time.Now().Sub(k.Created) > k.TTL { + if err := b.deleteKey(path, key); err != nil { + return nil, err + } + return nil, &backend.NotFoundError{ + Message: fmt.Sprintf("%v: %v not found", path, key)} + } + return k.Value, nil +} + +func (b *BoltBackend) DeleteKey(path []string, key string) error { + return b.deleteKey(path, key) +} + +func (b *BoltBackend) DeleteBucket(path []string, bucket string) error { + return b.deleteBucket(path, bucket) +} + func (b *BoltBackend) AcquireLock(token string, ttl time.Duration) error { b.Lock() defer b.Unlock() @@ -240,7 +304,7 @@ func (b *BoltBackend) DeleteUserKey(user, keyID string) error { } func (b *BoltBackend) UpsertServer(s backend.Server, ttl time.Duration) error { - return b.upsertJSONKey([]string{"servers"}, "val", s) + return b.upsertJSONKey([]string{"servers"}, s.ID, s) } func (b *BoltBackend) GetServers() ([]backend.Server, error) { @@ -444,6 +508,25 @@ func (b *BoltBackend) getKey(buckets []string, key string, val *[]byte) error { }) } +func (b *BoltBackend) getKeys(buckets []string) ([]string, error) { + out := []string{} + err := b.db.View(func(tx *bolt.Tx) error { + bkt, err := getBucket(tx, buckets) + if err != nil { + return err + } + c := bkt.Cursor() + for k, _ := c.First(); k != nil; k, _ = c.Next() { + out = append(out, string(k)) + } + return nil + }) + if err != nil { + return nil, err + } + return out, nil +} + func upsertBucket(b *bolt.Tx, buckets []string) (*bolt.Bucket, error) { bkt, err := b.CreateBucketIfNotExists([]byte(buckets[0])) if err != nil { @@ -478,3 +561,9 @@ func isNotFound(err error) bool { _, ok := err.(*backend.NotFoundError) return ok } + +type kv struct { + Created time.Time `json:"created"` + TTL time.Duration `json:"ttl"` + Value []byte `json:"val"` +} diff --git a/backend/boltbk/boltbk_test.go b/backend/boltbk/boltbk_test.go index bad6196ccf1..f0d341bbf8c 100644 --- a/backend/boltbk/boltbk_test.go +++ b/backend/boltbk/boltbk_test.go @@ -77,3 +77,11 @@ func (s *BoltSuite) TestToken(c *C) { func (s *BoltSuite) TestRemoteCert(c *C) { s.suite.RemoteCertCRUD(c) } + +func (s *BoltSuite) TestBasicCRUD(c *C) { + s.suite.BasicCRUD(c) +} + +func (s *BoltSuite) TestExpiration(c *C) { + s.suite.Expiration(c) +} diff --git a/backend/codec.go b/backend/codec.go new file mode 100644 index 00000000000..8c4216eb9db --- /dev/null +++ b/backend/codec.go @@ -0,0 +1,30 @@ +package backend + +import ( + "encoding/json" + "time" +) + +type JSONCodec struct { + Backend +} + +func (c *JSONCodec) UpsertJSONVal(path []string, key string, val interface{}, ttl time.Duration) error { + bytes, err := json.Marshal(val) + if err != nil { + return err + } + return c.UpsertVal(path, key, bytes, ttl) +} + +func (c *JSONCodec) GetJSONVal(path []string, key string, val interface{}) error { + bytes, err := json.Marshal(val) + if err != nil { + return err + } + bytes, err = c.GetVal(path, key) + if err != nil { + return err + } + return json.Unmarshal(bytes, val) +} diff --git a/backend/etcdbk/etcd.go b/backend/etcdbk/etcd.go index 7d1fac97f64..23f3b8231aa 100644 --- a/backend/etcdbk/etcd.go +++ b/backend/etcdbk/etcd.go @@ -4,6 +4,7 @@ package etcdbk import ( "encoding/json" "fmt" + "sort" "strings" "time" @@ -70,6 +71,39 @@ func (b *bk) reconnect() error { return nil } +func (b *bk) GetKeys(path []string) ([]string, error) { + keys, err := b.getKeys(b.key(path...)) + if err != nil { + return nil, err + } + sort.Sort(sort.StringSlice(keys)) + return keys, nil +} + +func (b *bk) UpsertVal(path []string, key string, val []byte, ttl time.Duration) error { + _, err := b.client.Set( + b.key(append(path, key)...), string(val), uint64(ttl/time.Second)) + return convertErr(err) +} + +func (b *bk) GetVal(path []string, key string) ([]byte, error) { + re, err := b.client.Get(b.key(append(path, key)...), false, false) + if err != nil { + return nil, convertErr(err) + } + return []byte(re.Node.Value), nil +} + +func (b *bk) DeleteKey(path []string, key string) error { + _, err := b.client.Delete(b.key(append(path, key)...), false) + return convertErr(err) +} + +func (b *bk) DeleteBucket(path []string, key string) error { + _, err := b.client.Delete(b.key(append(path, key)...), false) + return convertErr(err) +} + func (b *bk) AcquireLock(token string, ttl time.Duration) error { _, err := b.client.Create( b.key("locks", token), "lock", uint64(ttl/time.Second)) diff --git a/backend/etcdbk/etcd_test.go b/backend/etcdbk/etcd_test.go index e421832c5a2..c0f9d6c7458 100644 --- a/backend/etcdbk/etcd_test.go +++ b/backend/etcdbk/etcd_test.go @@ -121,3 +121,11 @@ func (s *EtcdSuite) TestToken(c *C) { func (s *EtcdSuite) TestRemoteCert(c *C) { s.suite.RemoteCertCRUD(c) } + +func (s *EtcdSuite) TestBasicCRUD(c *C) { + s.suite.BasicCRUD(c) +} + +func (s *EtcdSuite) TestExpiration(c *C) { + s.suite.Expiration(c) +} diff --git a/backend/membk/mem.go b/backend/membk/mem.go index 3793099c3d5..2dee68be774 100644 --- a/backend/membk/mem.go +++ b/backend/membk/mem.go @@ -40,6 +40,26 @@ func New() *MemBackend { } } +func (b *MemBackend) GetKeys(path []string) ([]string, error) { + return nil, nil +} + +func (b *MemBackend) UpsertVal(path []string, key string, val []byte, ttl time.Duration) error { + return nil +} + +func (b *MemBackend) GetVal(path []string, key string) ([]byte, error) { + return nil, nil +} + +func (b *MemBackend) DeleteKey(path []string, key string) error { + return nil +} + +func (b *MemBackend) DeleteBucket(path []string, key string) error { + return nil +} + func (b *MemBackend) AcquireLock(token string, ttl time.Duration) error { b.Lock() defer b.Unlock() diff --git a/backend/test/suite.go b/backend/test/suite.go index 418ef5f0243..f8f88b47d55 100644 --- a/backend/test/suite.go +++ b/backend/test/suite.go @@ -291,6 +291,42 @@ func (s *BackendSuite) RemoteCertCRUD(c *C) { c.Assert(err, FitsTypeOf, &backend.NotFoundError{}) } +func (s *BackendSuite) BasicCRUD(c *C) { + keys, err := s.B.GetKeys([]string{"keys"}) + c.Assert(err, IsNil) + c.Assert(keys, DeepEquals, []string{}) + + c.Assert(s.B.UpsertVal([]string{"a", "b"}, "bkey", []byte("val1"), 0), IsNil) + c.Assert(s.B.UpsertVal([]string{"a", "b"}, "akey", []byte("val2"), 0), IsNil) + + keys, err = s.B.GetKeys([]string{"a", "b"}) + c.Assert(err, IsNil) + c.Assert(keys, DeepEquals, []string{"akey", "bkey"}) + + out, err := s.B.GetVal([]string{"a", "b"}, "bkey") + c.Assert(err, IsNil) + c.Assert(string(out), Equals, "val1") + + c.Assert(s.B.UpsertVal([]string{"a", "b"}, "bkey", []byte("val-updated"), 0), IsNil) + out, err = s.B.GetVal([]string{"a", "b"}, "bkey") + c.Assert(err, IsNil) + c.Assert(string(out), Equals, "val-updated") + + c.Assert(s.B.DeleteKey([]string{"a", "b"}, "bkey"), IsNil) + c.Assert(s.B.DeleteKey([]string{"a", "b"}, "bkey"), FitsTypeOf, &backend.NotFoundError{}) +} + +func (s *BackendSuite) Expiration(c *C) { + c.Assert(s.B.UpsertVal([]string{"a", "b"}, "bkey", []byte("val1"), time.Second), IsNil) + c.Assert(s.B.UpsertVal([]string{"a", "b"}, "akey", []byte("val2"), 0), IsNil) + + time.Sleep(2 * time.Second) + + keys, err := s.B.GetKeys([]string{"a", "b"}) + c.Assert(err, IsNil) + c.Assert(keys, DeepEquals, []string{"akey"}) +} + func toSet(vals []string) map[string]struct{} { out := make(map[string]struct{}, len(vals)) for _, v := range vals { diff --git a/session/session.go b/session/session.go new file mode 100644 index 00000000000..982fe1f960a --- /dev/null +++ b/session/session.go @@ -0,0 +1,89 @@ +package session + +import ( + "time" + + "github.com/gravitational/teleport/backend" +) + +type SessionServer interface { + GetSessions() ([]Session, error) + GetSession(id string) (*Session, error) + DeleteSession(id string) error + UpsertParty(id string, p Party, ttl time.Duration) error +} + +type server struct { + bk backend.JSONCodec +} + +func New(bk backend.Backend) *server { + return &server{ + bk: backend.JSONCodec{bk}, + } +} + +func (s *server) GetSessions() ([]Session, error) { + keys, err := s.bk.GetKeys([]string{"sessions"}) + if err != nil { + return nil, err + } + out := []Session{} + for _, sid := range keys { + se, err := s.GetSession(sid) + if backend.IsNotFound(err) { + continue + } + out = append(out, *se) + } + return out, nil +} + +func (s *server) GetSession(id string) (*Session, error) { + if _, err := s.bk.GetVal([]string{"sessions", id}, "val"); err != nil { + return nil, err + } + + parties, err := s.bk.GetKeys([]string{"sessions", id, "parties"}) + + if err != nil { + return nil, err + } + out := []Party{} + for _, pk := range parties { + var p *Party + err := s.bk.GetJSONVal([]string{"sessions", id, "parties"}, pk, &p) + if err != nil { + if backend.IsNotFound(err) { // key was expired + continue + } + return nil, err + } + out = append(out, *p) + } + return &Session{ID: id, Parties: out}, nil +} + +func (s *server) UpsertParty(id string, p Party, ttl time.Duration) error { + if err := s.bk.UpsertVal([]string{"sessions", id}, "val", []byte("val"), ttl); err != nil { + return err + } + return s.bk.UpsertJSONVal([]string{"sessions", id, "parties"}, p.ID, p, ttl) +} + +func (s *server) DeleteSession(id string) error { + return s.bk.DeleteBucket([]string{"sessions"}, id) +} + +type Session struct { + ID string `json:"id"` + Parties []Party +} + +type Party struct { + ID string `json:"id"` + Site string `json:"site"` + User string `json:"user"` + Server string `json:"server"` + LastActive time.Time `json:"last_active"` +} diff --git a/session/session_test.go b/session/session_test.go new file mode 100644 index 00000000000..9aea5fc7e02 --- /dev/null +++ b/session/session_test.go @@ -0,0 +1,95 @@ +package session + +import ( + "path/filepath" + "testing" + "time" + + "github.com/gravitational/teleport/backend" + "github.com/gravitational/teleport/backend/boltbk" + + . "github.com/gravitational/teleport/Godeps/_workspace/src/gopkg.in/check.v1" +) + +func TestSessions(t *testing.T) { TestingT(t) } + +type BoltSuite struct { + bk *boltbk.BoltBackend + dir string + srv SessionServer +} + +var _ = Suite(&BoltSuite{}) + +func (s *BoltSuite) SetUpTest(c *C) { + s.dir = c.MkDir() + + var err error + s.bk, err = boltbk.New(filepath.Join(s.dir, "db")) + c.Assert(err, IsNil) + + s.srv = New(s.bk) +} + +func (s *BoltSuite) TearDownTest(c *C) { + c.Assert(s.bk.Close(), IsNil) +} + +func (s *BoltSuite) TestSessionsCRUD(c *C) { + out, err := s.srv.GetSessions() + c.Assert(err, IsNil) + c.Assert(out, DeepEquals, []Session{}) + + p1 := Party{ + ID: "p1", + User: "bob", + Site: "example.com", + Server: "localhost:1", + LastActive: time.Date(2009, time.November, 10, 23, 0, 0, 0, time.UTC), + } + c.Assert(s.srv.UpsertParty("s1", p1, 0), IsNil) + + out, err = s.srv.GetSessions() + c.Assert(err, IsNil) + sess := Session{ + ID: "s1", + Parties: []Party{p1}, + } + c.Assert(out, DeepEquals, []Session{sess}) + + // add one more party + p2 := Party{ + ID: "p2", + User: "alice", + Site: "example.com", + Server: "localhost:2", + LastActive: time.Date(2009, time.November, 10, 23, 1, 0, 0, time.UTC), + } + c.Assert(s.srv.UpsertParty("s1", p2, 0), IsNil) + + out, err = s.srv.GetSessions() + c.Assert(err, IsNil) + sess = Session{ + ID: "s1", + Parties: []Party{p1, p2}, + } + c.Assert(out, DeepEquals, []Session{sess}) + + // Update session party + p1.LastActive = time.Date(2009, time.November, 10, 23, 4, 0, 0, time.UTC) + c.Assert(s.srv.UpsertParty("s1", p1, 0), IsNil) + out, err = s.srv.GetSessions() + c.Assert(err, IsNil) + sess = Session{ + ID: "s1", + Parties: []Party{p1, p2}, + } + c.Assert(out, DeepEquals, []Session{sess}) + + // Delete session + c.Assert(s.srv.DeleteSession("s1"), IsNil) + c.Assert(s.srv.DeleteSession("s1"), FitsTypeOf, &backend.NotFoundError{}) + + _, err = s.srv.GetSession("s1") + c.Assert(err, FitsTypeOf, &backend.NotFoundError{}) +}