Răsfoiți Sursa

fix TOC profile lookup

Mike 5 luni în urmă
părinte
comite
e29649d886

+ 3 - 0
.mockery.yaml

@@ -204,6 +204,9 @@ packages:
       CookieBaker:
         config:
           filename: "mock_cookie_baker_test.go"
+      SessionRetriever:
+        config:
+          filename: "mock_session_retriever_test.go"
   github.com/mk6i/open-oscar-server/server/kerberos:
     interfaces:
       AuthService:

+ 1 - 0
cmd/server/factory.go

@@ -453,6 +453,7 @@ func TOC(deps Container) *toc.Server {
 			ChatNavService:    foodgroup.NewChatNavService(logger, deps.sqLiteUserStore),
 			SNACRateLimits:    deps.snacRateLimits,
 			HTTPIPRateLimiter: toc.NewIPRateLimiter(rate.Every(1*time.Minute), 10, 1*time.Minute),
+			SessionRetriever:  deps.inMemorySessionManager,
 		},
 		toc.NewIPRateLimiter(rate.Every(1*time.Minute), 10, 1*time.Minute),
 		deps.icbmSvc.RestoreWarningLevel,

+ 1 - 0
server/toc/cmd_client.go

@@ -129,6 +129,7 @@ type OSCARProxy struct {
 	OServiceService   OServiceService
 	PermitDenyService PermitDenyService
 	TOCConfigStore    TOCConfigStore
+	SessionRetriever  SessionRetriever
 	SNACRateLimits    wire.SNACRateLimits
 	HTTPIPRateLimiter *IPRateLimiter
 }

+ 10 - 0
server/toc/helpers_test.go

@@ -247,6 +247,15 @@ type userParams []struct {
 	err          error
 }
 
+type retrieveSessionParams []struct {
+	screenName      state.IdentScreenName
+	returnedSession *state.Session
+}
+
+type sessionRetrieverParams struct {
+	retrieveSessionParams
+}
+
 type tocConfigParams struct {
 	setTOCConfigParams
 	userParams
@@ -265,6 +274,7 @@ type mockParams struct {
 	locateParams
 	oServiceParams
 	permitDenyParams
+	sessionRetrieverParams
 	tocConfigParams
 }
 

+ 13 - 3
server/toc/http.go

@@ -151,15 +151,25 @@ func (s OSCARProxy) ProfileHandler(w http.ResponseWriter, r *http.Request) {
 		return
 	}
 
-	sess := state.NewSession().AddInstance()
-	sess.Session().SetIdentScreenName(state.NewIdentScreenName(from))
+	sess := s.SessionRetriever.RetrieveSession(state.NewIdentScreenName(from))
+	if sess == nil {
+		http.Error(w, "invalid session", http.StatusForbidden)
+		return
+	}
+
+	instances := sess.Instances()
+	if len(instances) == 0 {
+		http.Error(w, "invalid session", http.StatusForbidden)
+		return
+	}
+
 	inBody := wire.SNAC_0x02_0x05_LocateUserInfoQuery{
 		Type:       uint16(wire.LocateTypeSig),
 		ScreenName: user,
 	}
 
 	ctx := r.Context()
-	info, err := s.LocateService.UserInfoQuery(ctx, sess, wire.SNACFrame{}, inBody)
+	info, err := s.LocateService.UserInfoQuery(ctx, instances[0], wire.SNACFrame{}, inBody)
 	if err != nil {
 		s.logAndReturn500(ctx, w, fmt.Errorf("LocateService.UserInfoQuery: %w", err))
 		return

+ 55 - 0
server/toc/http_test.go

@@ -72,6 +72,14 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 						},
 					},
 				},
+				sessionRetrieverParams: sessionRetrieverParams{
+					retrieveSessionParams: retrieveSessionParams{
+						{
+							screenName:      state.NewIdentScreenName("me"),
+							returnedSession: newTestSession("me").Session(),
+						},
+					},
+				},
 			},
 		},
 		{
@@ -105,6 +113,14 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 						},
 					},
 				},
+				sessionRetrieverParams: sessionRetrieverParams{
+					retrieveSessionParams: retrieveSessionParams{
+						{
+							screenName:      state.NewIdentScreenName("me"),
+							returnedSession: newTestSession("me").Session(),
+						},
+					},
+				},
 			},
 		},
 		{
@@ -165,6 +181,14 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 						},
 					},
 				},
+				sessionRetrieverParams: sessionRetrieverParams{
+					retrieveSessionParams: retrieveSessionParams{
+						{
+							screenName:      state.NewIdentScreenName("me"),
+							returnedSession: newTestSession("me").Session(),
+						},
+					},
+				},
 			},
 		},
 		{
@@ -191,6 +215,14 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 						},
 					},
 				},
+				sessionRetrieverParams: sessionRetrieverParams{
+					retrieveSessionParams: retrieveSessionParams{
+						{
+							screenName:      state.NewIdentScreenName("me"),
+							returnedSession: newTestSession("me").Session(),
+						},
+					},
+				},
 			},
 		},
 		{
@@ -219,6 +251,14 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 						},
 					},
 				},
+				sessionRetrieverParams: sessionRetrieverParams{
+					retrieveSessionParams: retrieveSessionParams{
+						{
+							screenName:      state.NewIdentScreenName("me"),
+							returnedSession: newTestSession("me").Session(),
+						},
+					},
+				},
 			},
 		},
 		{
@@ -247,6 +287,14 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 						},
 					},
 				},
+				sessionRetrieverParams: sessionRetrieverParams{
+					retrieveSessionParams: retrieveSessionParams{
+						{
+							screenName:      state.NewIdentScreenName("me"),
+							returnedSession: newTestSession("me").Session(),
+						},
+					},
+				},
 			},
 		},
 		{
@@ -622,12 +670,19 @@ func TestOSCARProxy_NewServeMux(t *testing.T) {
 					InfoQuery(mock.Anything, wire.SNACFrame{}, params.inBody).
 					Return(params.msg, params.err)
 			}
+			sessionRetriever := newMockSessionRetriever(t)
+			for _, params := range tc.mockParams.retrieveSessionParams {
+				sessionRetriever.EXPECT().
+					RetrieveSession(params.screenName).
+					Return(params.returnedSession)
+			}
 
 			svc := OSCARProxy{
 				CookieBaker:       cookieBaker,
 				DirSearchService:  dirSearchSvc,
 				LocateService:     locateSvc,
 				Logger:            slog.Default(),
+				SessionRetriever:  sessionRetriever,
 				HTTPIPRateLimiter: NewIPRateLimiter(rate.Every(1*time.Minute), 10, 1*time.Minute),
 				SNACRateLimits:    wire.DefaultSNACRateLimits(),
 			}

+ 90 - 0
server/toc/mock_session_retriever_test.go

@@ -0,0 +1,90 @@
+// Code generated by mockery; DO NOT EDIT.
+// github.com/vektra/mockery
+// template: testify
+
+package toc
+
+import (
+	"github.com/mk6i/open-oscar-server/state"
+	mock "github.com/stretchr/testify/mock"
+)
+
+// newMockSessionRetriever creates a new instance of mockSessionRetriever. It also registers a testing interface on the mock and a cleanup function to assert the mocks expectations.
+// The first argument is typically a *testing.T value.
+func newMockSessionRetriever(t interface {
+	mock.TestingT
+	Cleanup(func())
+}) *mockSessionRetriever {
+	mock := &mockSessionRetriever{}
+	mock.Mock.Test(t)
+
+	t.Cleanup(func() { mock.AssertExpectations(t) })
+
+	return mock
+}
+
+// mockSessionRetriever is an autogenerated mock type for the SessionRetriever type
+type mockSessionRetriever struct {
+	mock.Mock
+}
+
+type mockSessionRetriever_Expecter struct {
+	mock *mock.Mock
+}
+
+func (_m *mockSessionRetriever) EXPECT() *mockSessionRetriever_Expecter {
+	return &mockSessionRetriever_Expecter{mock: &_m.Mock}
+}
+
+// RetrieveSession provides a mock function for the type mockSessionRetriever
+func (_mock *mockSessionRetriever) RetrieveSession(screenName state.IdentScreenName) *state.Session {
+	ret := _mock.Called(screenName)
+
+	if len(ret) == 0 {
+		panic("no return value specified for RetrieveSession")
+	}
+
+	var r0 *state.Session
+	if returnFunc, ok := ret.Get(0).(func(state.IdentScreenName) *state.Session); ok {
+		r0 = returnFunc(screenName)
+	} else {
+		if ret.Get(0) != nil {
+			r0 = ret.Get(0).(*state.Session)
+		}
+	}
+	return r0
+}
+
+// mockSessionRetriever_RetrieveSession_Call is a *mock.Call that shadows Run/Return methods with type explicit version for method 'RetrieveSession'
+type mockSessionRetriever_RetrieveSession_Call struct {
+	*mock.Call
+}
+
+// RetrieveSession is a helper method to define mock.On call
+//   - screenName state.IdentScreenName
+func (_e *mockSessionRetriever_Expecter) RetrieveSession(screenName interface{}) *mockSessionRetriever_RetrieveSession_Call {
+	return &mockSessionRetriever_RetrieveSession_Call{Call: _e.mock.On("RetrieveSession", screenName)}
+}
+
+func (_c *mockSessionRetriever_RetrieveSession_Call) Run(run func(screenName state.IdentScreenName)) *mockSessionRetriever_RetrieveSession_Call {
+	_c.Call.Run(func(args mock.Arguments) {
+		var arg0 state.IdentScreenName
+		if args[0] != nil {
+			arg0 = args[0].(state.IdentScreenName)
+		}
+		run(
+			arg0,
+		)
+	})
+	return _c
+}
+
+func (_c *mockSessionRetriever_RetrieveSession_Call) Return(session *state.Session) *mockSessionRetriever_RetrieveSession_Call {
+	_c.Call.Return(session)
+	return _c
+}
+
+func (_c *mockSessionRetriever_RetrieveSession_Call) RunAndReturn(run func(screenName state.IdentScreenName) *state.Session) *mockSessionRetriever_RetrieveSession_Call {
+	_c.Call.Return(run)
+	return _c
+}

+ 1 - 1
server/toc/server.go

@@ -76,7 +76,7 @@ func (l *channelListener) Accept() (net.Conn, error) {
 	case <-l.ctx.Done():
 		return nil, io.EOF
 	case ch := <-l.ch:
-		return ch, io.EOF
+		return ch, nil
 	}
 }
 

+ 9 - 0
server/toc/types.go

@@ -104,3 +104,12 @@ type CookieBaker interface {
 type AdminService interface {
 	InfoChangeRequest(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x07_0x04_AdminInfoChangeRequest) (wire.SNACMessage, error)
 }
+
+// SessionRetriever defines a method for retrieving an active session
+// associated with a given screen name.
+type SessionRetriever interface {
+	// RetrieveSession returns the session associated with the given screen name,
+	// or nil if no active session exists. Returns the Session object if there
+	// are active instances with complete signon.
+	RetrieveSession(screenName state.IdentScreenName) *state.Session
+}