Browse Source

client thread safety

Jonathan Turner 8 years ago
parent
commit
17b157271d
6 changed files with 96 additions and 18 deletions
  1. 1 1
      .travis.yml
  2. 8 0
      client/cache.go
  3. 19 10
      client/client.go
  4. 50 0
      client/client_integration_test.go
  5. 12 4
      client/session.go
  6. 6 3
      messages/KDCReq.go

+ 1 - 1
.travis.yml

@@ -8,7 +8,7 @@ go:
 
 go_import_path: gopkg.in/jcmturner/gokrb5.v1
 
-gobuild_args: -tags=integration
+gobuild_args: -tags=integration -race
 
 sudo: required
 

+ 8 - 0
client/cache.go

@@ -4,12 +4,14 @@ import (
 	"gopkg.in/jcmturner/gokrb5.v1/messages"
 	"gopkg.in/jcmturner/gokrb5.v1/types"
 	"strings"
+	"sync"
 	"time"
 )
 
 // Cache for client tickets.
 type Cache struct {
 	Entries map[string]CacheEntry
+	mux     sync.RWMutex
 }
 
 // CacheEntry holds details for a client cache entry.
@@ -31,6 +33,8 @@ func NewCache() *Cache {
 
 // GetEntry returns a cache entry that matches the SPN.
 func (c *Cache) getEntry(spn string) (CacheEntry, bool) {
+	c.mux.RLock()
+	defer c.mux.RUnlock()
 	e, ok := (*c).Entries[spn]
 	return e, ok
 }
@@ -38,6 +42,8 @@ func (c *Cache) getEntry(spn string) (CacheEntry, bool) {
 // AddEntry adds a ticket to the cache.
 func (c *Cache) addEntry(tkt messages.Ticket, authTime, startTime, endTime, renewTill time.Time, sessionKey types.EncryptionKey) CacheEntry {
 	spn := strings.Join(tkt.SName.NameString, "/")
+	c.mux.Lock()
+	defer c.mux.Unlock()
 	(*c).Entries[spn] = CacheEntry{
 		Ticket:     tkt,
 		AuthTime:   authTime,
@@ -51,6 +57,8 @@ func (c *Cache) addEntry(tkt messages.Ticket, authTime, startTime, endTime, rene
 
 // RemoveEntry removes the cache entry for the defined SPN.
 func (c *Cache) RemoveEntry(spn string) {
+	c.mux.Lock()
+	defer c.mux.Unlock()
 	delete(c.Entries, spn)
 }
 

+ 19 - 10
client/client.go

@@ -20,7 +20,7 @@ type Client struct {
 	Credentials *credentials.Credentials
 	Config      *config.Config
 	GoKrb5Conf  *Config
-	sessions    sessions
+	sessions    *sessions
 	Cache       *Cache
 }
 
@@ -39,8 +39,10 @@ func NewClientWithPassword(username, realm, password string) Client {
 		Credentials: creds.WithPassword(password),
 		Config:      config.NewConfig(),
 		GoKrb5Conf:  &Config{},
-		sessions:    make(sessions),
-		Cache:       NewCache(),
+		sessions: &sessions{
+			Entries: make(map[string]*session),
+		},
+		Cache: NewCache(),
 	}
 }
 
@@ -51,8 +53,10 @@ func NewClientWithKeytab(username, realm string, kt keytab.Keytab) Client {
 		Credentials: creds.WithKeytab(kt),
 		Config:      config.NewConfig(),
 		GoKrb5Conf:  &Config{},
-		sessions:    make(sessions),
-		Cache:       NewCache(),
+		sessions: &sessions{
+			Entries: make(map[string]*session),
+		},
+		Cache: NewCache(),
 	}
 }
 
@@ -64,8 +68,10 @@ func NewClientFromCCache(c credentials.CCache) (Client, error) {
 		Credentials: c.GetClientCredentials(),
 		Config:      config.NewConfig(),
 		GoKrb5Conf:  &Config{},
-		sessions:    make(sessions),
-		Cache:       NewCache(),
+		sessions: &sessions{
+			Entries: make(map[string]*session),
+		},
+		Cache: NewCache(),
 	}
 	spn := types.PrincipalName{
 		NameType:   nametype.KRB_NT_SRV_INST,
@@ -80,7 +86,7 @@ func NewClientFromCCache(c credentials.CCache) (Client, error) {
 	if err != nil {
 		return cl, fmt.Errorf("TGT bytes in cache are not valid: %v", err)
 	}
-	cl.sessions[c.DefaultPrincipal.Realm] = &session{
+	cl.sessions.Entries[c.DefaultPrincipal.Realm] = &session{
 		Realm:      c.DefaultPrincipal.Realm,
 		AuthTime:   cred.AuthTime,
 		EndTime:    cred.EndTime,
@@ -159,8 +165,11 @@ func (cl *Client) LoadConfig(cfgPath string) (*Client, error) {
 // IsConfigured indicates if the client has the values required set.
 func (cl *Client) IsConfigured() (bool, error) {
 	// Client needs to have either a password, keytab or a session already (later when loading from CCache)
-	if !cl.Credentials.HasPassword() && !cl.Credentials.HasKeytab() && cl.sessions[cl.Config.LibDefaults.DefaultRealm].AuthTime.IsZero() {
-		return false, errors.New("client has neither a keytab nor a password set and no session")
+	if !cl.Credentials.HasPassword() && !cl.Credentials.HasKeytab() {
+		sess, err := cl.GetSessionFromRealm(cl.Config.LibDefaults.DefaultRealm)
+		if err != nil || sess.AuthTime.IsZero() {
+			return false, errors.New("client has neither a keytab nor a password set and no session")
+		}
 	}
 	if cl.Credentials.Username == "" {
 		return false, errors.New("client does not have a username")

+ 50 - 0
client/client_integration_test.go

@@ -262,6 +262,56 @@ func TestClient_SetSPNEGOHeader(t *testing.T) {
 	assert.Equal(t, http.StatusOK, httpResp.StatusCode, "Status code in response to client SPNEGO request not as expected")
 }
 
+func TestMultiThreadedClientUse(t *testing.T) {
+	b, _ := hex.DecodeString(testdata.TESTUSER1_KEYTAB)
+	kt, _ := keytab.Parse(b)
+	c, _ := config.NewConfigFromString(testdata.TEST_KRB5CONF)
+	addr := os.Getenv("TEST_KDC_ADDR")
+	if addr == "" {
+		addr = testdata.TEST_KDC_ADDR
+	}
+	c.Realms[0].KDC = []string{addr + ":" + testdata.TEST_KDC}
+	cl := NewClientWithKeytab("testuser1", "TEST.GOKRB5", kt)
+	cl.WithConfig(c)
+
+	for i := 0; i < 5; i++ {
+		go login(t, &cl)
+	}
+
+	for i := 0; i < 5; i++ {
+		go spnegoGet(t, &cl)
+	}
+}
+
+func login(t *testing.T, cl *Client) {
+	err := cl.Login()
+	if err != nil {
+		t.Fatalf("Error on AS_REQ: %v\n", err)
+	}
+}
+
+func spnegoGet(t *testing.T, cl *Client) {
+	url := os.Getenv("TEST_HTTP_URL")
+	if url == "" {
+		url = testdata.TEST_HTTP_URL
+	}
+	r, _ := http.NewRequest("GET", url, nil)
+	httpResp, err := http.DefaultClient.Do(r)
+	if err != nil {
+		t.Fatalf("Request error: %v\n", err)
+	}
+	assert.Equal(t, http.StatusUnauthorized, httpResp.StatusCode, "Status code in response to client with no SPNEGO not as expected")
+	err = cl.SetSPNEGOHeader(r, "HTTP/host.test.gokrb5")
+	if err != nil {
+		t.Fatalf("Error setting client SPNEGO header: %v", err)
+	}
+	httpResp, err = http.DefaultClient.Do(r)
+	if err != nil {
+		t.Fatalf("Request error: %v\n", err)
+	}
+	assert.Equal(t, http.StatusOK, httpResp.StatusCode, "Status code in response to client SPNEGO request not as expected")
+}
+
 func TestNewClientFromCCache(t *testing.T) {
 	b, err := hex.DecodeString(testdata.CCACHE_TEST)
 	if err != nil {

+ 12 - 4
client/session.go

@@ -6,11 +6,15 @@ import (
 	"gopkg.in/jcmturner/gokrb5.v1/krberror"
 	"gopkg.in/jcmturner/gokrb5.v1/messages"
 	"gopkg.in/jcmturner/gokrb5.v1/types"
+	"sync"
 	"time"
 )
 
 // Sessions keyed on the realm name
-type sessions map[string]*session
+type sessions struct {
+	Entries map[string]*session
+	mux     sync.RWMutex
+}
 
 // Client session struct.
 type session struct {
@@ -25,6 +29,8 @@ type session struct {
 
 //
 func (cl *Client) AddSession(tkt messages.Ticket, dep messages.EncKDCRepPart) {
+	cl.sessions.mux.Lock()
+	defer cl.sessions.mux.Unlock()
 	s := &session{
 		Realm:                tkt.SName.NameString[1],
 		AuthTime:             dep.AuthTime,
@@ -34,7 +40,7 @@ func (cl *Client) AddSession(tkt messages.Ticket, dep messages.EncKDCRepPart) {
 		SessionKey:           dep.Key,
 		SessionKeyExpiration: dep.KeyExpiration,
 	}
-	cl.sessions[tkt.SName.NameString[1]] = s
+	cl.sessions.Entries[tkt.SName.NameString[1]] = s
 	cl.EnableAutoSessionRenewal(s)
 }
 
@@ -91,9 +97,11 @@ func (cl *Client) updateSession(s *session) error {
 
 func (cl *Client) GetSessionFromRealm(realm string) (sess *session, err error) {
 	var ok bool
-	sess, ok = cl.sessions[realm]
+	cl.sessions.mux.RLock()
+	defer cl.sessions.mux.RUnlock()
+	sess, ok = cl.sessions.Entries[realm]
 	if !ok {
-		sess, ok = cl.sessions[cl.Config.LibDefaults.DefaultRealm]
+		sess, ok = cl.sessions.Entries[cl.Config.LibDefaults.DefaultRealm]
 		if !ok {
 			err = fmt.Errorf("client does not have a session for realm %s or for the default realm %s, login first", realm, cl.Config.LibDefaults.DefaultRealm)
 			return

+ 6 - 3
messages/KDCReq.go

@@ -89,21 +89,24 @@ func NewASReq(realm string, c *config.Config, cname types.PrincipalName) (ASReq,
 		return ASReq{}, err
 	}
 	t := time.Now().UTC()
+	// Copy the default options to make this thread safe
+	var kopts asn1.BitString
+	copy(kopts.Bytes, c.LibDefaults.KDCDefaultOptions.Bytes)
+	kopts.BitLength = c.LibDefaults.KDCDefaultOptions.BitLength
 	a := ASReq{
 		KDCReqFields{
 			PVNO:    iana.PVNO,
 			MsgType: msgtype.KRB_AS_REQ,
 			PAData:  types.PADataSequence{},
 			ReqBody: KDCReqBody{
-				KDCOptions: c.LibDefaults.KDCDefaultOptions,
+				KDCOptions: kopts,
 				Realm:      realm,
 				CName:      cname,
 				SName: types.PrincipalName{
 					NameType:   nametype.KRB_NT_SRV_INST,
 					NameString: []string{"krbtgt", realm},
 				},
-				Till: t.Add(c.LibDefaults.TicketLifetime),
-				//Till:  t.Add(time.Duration(24) * time.Hour),
+				Till:  t.Add(c.LibDefaults.TicketLifetime),
 				Nonce: int(nonce.Int64()),
 				EType: c.LibDefaults.DefaultTktEnctypeIDs,
 			},