Procházet zdrojové kódy

Merge pull request #163 from siohaza/os-motd

feat: add snac motd
Mike před 6 měsíci
rodič
revize
9ff8d863c4

+ 29 - 10
foodgroup/oservice.go

@@ -67,14 +67,16 @@ func NewOServiceService(
 // attempt to accommodate any particular food group version. The server
 // attempt to accommodate any particular food group version. The server
 // implicitly accommodates any food group version for Windows AIM clients 5.x.
 // implicitly accommodates any food group version for Windows AIM clients 5.x.
 // It returns SNAC wire.OServiceHostVersions containing the server's supported
 // It returns SNAC wire.OServiceHostVersions containing the server's supported
-// food group versions.
+// food group versions followed by SNAC wire.OServiceMotd containing Message of
+// the Day. MOTD is sent here because some clients such as Jimm wait for it
+// before sending RateParamsQuery, causing the login flow to stall if omitted.
 // todo this documentation
 // todo this documentation
-func (s OServiceService) ClientVersions(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) wire.SNACMessage {
+func (s OServiceService) ClientVersions(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) []wire.SNACMessage {
 	var versions [wire.MDir + 1]uint16
 	var versions [wire.MDir + 1]uint16
 
 
 	if len(inBody.Versions)%2 != 0 {
 	if len(inBody.Versions)%2 != 0 {
 		s.logger.ErrorContext(ctx, "got uneven food group length")
 		s.logger.ErrorContext(ctx, "got uneven food group length")
-		return wire.SNACMessage{}
+		return nil
 	}
 	}
 
 
 	for i := 0; i < len(inBody.Versions); i += 2 {
 	for i := 0; i < len(inBody.Versions); i += 2 {
@@ -93,14 +95,31 @@ func (s OServiceService) ClientVersions(ctx context.Context, instance *state.Ses
 
 
 	instance.SetFoodGroupVersions(versions)
 	instance.SetFoodGroupVersions(versions)
 
 
-	return wire.SNACMessage{
-		Frame: wire.SNACFrame{
-			FoodGroup: wire.OService,
-			SubGroup:  wire.OServiceHostVersions,
-			RequestID: inFrame.RequestID,
+	return []wire.SNACMessage{
+		{
+			Frame: wire.SNACFrame{
+				FoodGroup: wire.OService,
+				SubGroup:  wire.OServiceHostVersions,
+				RequestID: inFrame.RequestID,
+			},
+			Body: wire.SNAC_0x01_0x18_OServiceHostVersions{
+				Versions: inBody.Versions,
+			},
 		},
 		},
-		Body: wire.SNAC_0x01_0x18_OServiceHostVersions{
-			Versions: inBody.Versions,
+		{
+			Frame: wire.SNACFrame{
+				FoodGroup: wire.OService,
+				SubGroup:  wire.OServiceMotd,
+				RequestID: wire.ReqIDFromServer,
+			},
+			Body: wire.SNAC_0x01_0x13_OServiceMOTD{
+				MessageType: 0x0004,
+				TLVRestBlock: wire.TLVRestBlock{
+					TLVList: wire.TLVList{
+						wire.NewTLVBE(wire.OServiceTLVTagsMOTDMessage, "Welcome to Open OSCAR Server"),
+					},
+				},
+			},
 		},
 		},
 	}
 	}
 }
 }

+ 25 - 8
foodgroup/oservice_test.go

@@ -1773,14 +1773,31 @@ func TestOServiceService_ClientVersions(t *testing.T) {
 		logger: slog.Default(),
 		logger: slog.Default(),
 	}
 	}
 
 
-	want := wire.SNACMessage{
-		Frame: wire.SNACFrame{
-			FoodGroup: wire.OService,
-			SubGroup:  wire.OServiceHostVersions,
-			RequestID: 1234,
-		},
-		Body: wire.SNAC_0x01_0x18_OServiceHostVersions{
-			Versions: []uint16{5, 6, 7, 8},
+	want := []wire.SNACMessage{
+		{
+			Frame: wire.SNACFrame{
+				FoodGroup: wire.OService,
+				SubGroup:  wire.OServiceHostVersions,
+				RequestID: 1234,
+			},
+			Body: wire.SNAC_0x01_0x18_OServiceHostVersions{
+				Versions: []uint16{5, 6, 7, 8},
+			},
+		},
+		{
+			Frame: wire.SNACFrame{
+				FoodGroup: wire.OService,
+				SubGroup:  wire.OServiceMotd,
+				RequestID: wire.ReqIDFromServer,
+			},
+			Body: wire.SNAC_0x01_0x13_OServiceMOTD{
+				MessageType: 0x0004,
+				TLVRestBlock: wire.TLVRestBlock{
+					TLVList: wire.TLVList{
+						wire.NewTLVBE(wire.OServiceTLVTagsMOTDMessage, "Welcome to Open OSCAR Server"),
+					},
+				},
+			},
 		},
 		},
 	}
 	}
 
 

+ 8 - 3
server/oscar/handler.go

@@ -781,9 +781,14 @@ func (rt Handler) OServiceClientVersions(ctx context.Context, instance *state.Se
 	if err := wire.UnmarshalBE(&inBody, r); err != nil {
 	if err := wire.UnmarshalBE(&inBody, r); err != nil {
 		return err
 		return err
 	}
 	}
-	outSNAC := rt.OServiceService.ClientVersions(ctx, instance, inFrame, inBody)
-	rt.LogRequestAndResponse(ctx, inFrame, inBody, outSNAC.Frame, outSNAC.Body)
-	return rw.SendSNAC(outSNAC.Frame, outSNAC.Body)
+	outSNACs := rt.OServiceService.ClientVersions(ctx, instance, inFrame, inBody)
+	for _, snac := range outSNACs {
+		rt.LogRequestAndResponse(ctx, inFrame, inBody, snac.Frame, snac.Body)
+		if err := rw.SendSNAC(snac.Frame, snac.Body); err != nil {
+			return err
+		}
+	}
+	return nil
 }
 }
 
 
 func (rt Handler) OServiceSetUserInfoFields(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, r io.Reader, rw ResponseWriter) error {
 func (rt Handler) OServiceSetUserInfoFields(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, r io.Reader, rw ResponseWriter) error {

+ 35 - 10
server/oscar/handler_test.go

@@ -4254,14 +4254,31 @@ func TestHandler_OServiceServiceClientVersions(t *testing.T) {
 				},
 				},
 				Body: tt.inputBody,
 				Body: tt.inputBody,
 			}
 			}
-			output := wire.SNACMessage{
-				Frame: wire.SNACFrame{
-					FoodGroup: wire.OService,
-					SubGroup:  wire.OServiceHostVersions,
+			output := []wire.SNACMessage{
+				{
+					Frame: wire.SNACFrame{
+						FoodGroup: wire.OService,
+						SubGroup:  wire.OServiceHostVersions,
+					},
+					Body: wire.SNAC_0x01_0x18_OServiceHostVersions{
+						Versions: []uint16{
+							10,
+						},
+					},
 				},
 				},
-				Body: wire.SNAC_0x01_0x18_OServiceHostVersions{
-					Versions: []uint16{
-						10,
+				{
+					Frame: wire.SNACFrame{
+						FoodGroup: wire.OService,
+						SubGroup:  wire.OServiceMotd,
+						RequestID: wire.ReqIDFromServer,
+					},
+					Body: wire.SNAC_0x01_0x13_OServiceMOTD{
+						MessageType: 0x0004,
+						TLVRestBlock: wire.TLVRestBlock{
+							TLVList: wire.TLVList{
+								wire.NewTLVBE(wire.OServiceTLVTagsMOTDMessage, "Welcome to Open OSCAR Server"),
+							},
+						},
 					},
 					},
 				},
 				},
 			}
 			}
@@ -4280,9 +4297,17 @@ func TestHandler_OServiceServiceClientVersions(t *testing.T) {
 			}
 			}
 
 
 			responseWriter := newMockResponseWriter(t)
 			responseWriter := newMockResponseWriter(t)
-			responseWriter.EXPECT().
-				SendSNAC(output.Frame, output.Body).
-				Return(tt.responseError)
+			if tt.responseError == nil {
+				for _, snac := range output {
+					responseWriter.EXPECT().
+						SendSNAC(snac.Frame, snac.Body).
+						Return(nil)
+				}
+			} else {
+				responseWriter.EXPECT().
+					SendSNAC(output[0].Frame, output[0].Body).
+					Return(tt.responseError)
+			}
 
 
 			buf := &bytes.Buffer{}
 			buf := &bytes.Buffer{}
 			assert.NoError(t, wire.MarshalBE(input.Body, buf))
 			assert.NoError(t, wire.MarshalBE(input.Body, buf))

+ 8 - 8
server/oscar/mock_oservice_service_test.go

@@ -110,18 +110,18 @@ func (_c *mockOServiceService_ClientOnline_Call) RunAndReturn(run func(ctx conte
 }
 }
 
 
 // ClientVersions provides a mock function for the type mockOServiceService
 // ClientVersions provides a mock function for the type mockOServiceService
-func (_mock *mockOServiceService) ClientVersions(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) wire.SNACMessage {
+func (_mock *mockOServiceService) ClientVersions(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) []wire.SNACMessage {
 	ret := _mock.Called(ctx, instance, inFrame, inBody)
 	ret := _mock.Called(ctx, instance, inFrame, inBody)
 
 
 	if len(ret) == 0 {
 	if len(ret) == 0 {
 		panic("no return value specified for ClientVersions")
 		panic("no return value specified for ClientVersions")
 	}
 	}
 
 
-	var r0 wire.SNACMessage
-	if returnFunc, ok := ret.Get(0).(func(context.Context, *state.SessionInstance, wire.SNACFrame, wire.SNAC_0x01_0x17_OServiceClientVersions) wire.SNACMessage); ok {
+	var r0 []wire.SNACMessage
+	if returnFunc, ok := ret.Get(0).(func(context.Context, *state.SessionInstance, wire.SNACFrame, wire.SNAC_0x01_0x17_OServiceClientVersions) []wire.SNACMessage); ok {
 		r0 = returnFunc(ctx, instance, inFrame, inBody)
 		r0 = returnFunc(ctx, instance, inFrame, inBody)
-	} else {
-		r0 = ret.Get(0).(wire.SNACMessage)
+	} else if ret.Get(0) != nil {
+		r0 = ret.Get(0).([]wire.SNACMessage)
 	}
 	}
 	return r0
 	return r0
 }
 }
@@ -168,12 +168,12 @@ func (_c *mockOServiceService_ClientVersions_Call) Run(run func(ctx context.Cont
 	return _c
 	return _c
 }
 }
 
 
-func (_c *mockOServiceService_ClientVersions_Call) Return(sNACMessage wire.SNACMessage) *mockOServiceService_ClientVersions_Call {
-	_c.Call.Return(sNACMessage)
+func (_c *mockOServiceService_ClientVersions_Call) Return(sNACMessages []wire.SNACMessage) *mockOServiceService_ClientVersions_Call {
+	_c.Call.Return(sNACMessages)
 	return _c
 	return _c
 }
 }
 
 
-func (_c *mockOServiceService_ClientVersions_Call) RunAndReturn(run func(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) wire.SNACMessage) *mockOServiceService_ClientVersions_Call {
+func (_c *mockOServiceService_ClientVersions_Call) RunAndReturn(run func(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) []wire.SNACMessage) *mockOServiceService_ClientVersions_Call {
 	_c.Call.Return(run)
 	_c.Call.Return(run)
 	return _c
 	return _c
 }
 }

+ 1 - 1
server/oscar/types.go

@@ -150,7 +150,7 @@ type ODirService interface {
 
 
 type OServiceService interface {
 type OServiceService interface {
 	ClientOnline(ctx context.Context, service uint16, inBody wire.SNAC_0x01_0x02_OServiceClientOnline, instance *state.SessionInstance) error
 	ClientOnline(ctx context.Context, service uint16, inBody wire.SNAC_0x01_0x02_OServiceClientOnline, instance *state.SessionInstance) error
-	ClientVersions(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) wire.SNACMessage
+	ClientVersions(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame, inBody wire.SNAC_0x01_0x17_OServiceClientVersions) []wire.SNACMessage
 	HostOnline(service uint16) wire.SNACMessage
 	HostOnline(service uint16) wire.SNACMessage
 	IdleNotification(ctx context.Context, instance *state.SessionInstance, inBody wire.SNAC_0x01_0x11_OServiceIdleNotification) error
 	IdleNotification(ctx context.Context, instance *state.SessionInstance, inBody wire.SNAC_0x01_0x11_OServiceIdleNotification) error
 	RateParamsQuery(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame) wire.SNACMessage
 	RateParamsQuery(ctx context.Context, instance *state.SessionInstance, inFrame wire.SNACFrame) wire.SNACMessage

+ 6 - 0
wire/snacs.go

@@ -237,6 +237,7 @@ const (
 	OServiceTLVTagsSSLCertName   uint16 = 0x8D
 	OServiceTLVTagsSSLCertName   uint16 = 0x8D
 	OServiceTLVTagsSSLState      uint16 = 0x8E
 	OServiceTLVTagsSSLState      uint16 = 0x8E
 	OserviceTLVTagsSSLUseSSL     uint16 = 0x8C
 	OserviceTLVTagsSSLUseSSL     uint16 = 0x8C
+	OServiceTLVTagsMOTDMessage   uint16 = 0x0B
 
 
 	OServiceDiscErrNewLogin   uint8 = 0x01
 	OServiceDiscErrNewLogin   uint8 = 0x01
 	OServiceDiscErrAccDeleted uint8 = 0x02
 	OServiceDiscErrAccDeleted uint8 = 0x02
@@ -350,6 +351,11 @@ type SNAC_0x01_0x11_OServiceIdleNotification struct {
 	IdleTime uint32
 	IdleTime uint32
 }
 }
 
 
+type SNAC_0x01_0x13_OServiceMOTD struct {
+	MessageType uint16
+	TLVRestBlock
+}
+
 type SNAC_0x01_0x14_OServiceSetPrivacyFlags struct {
 type SNAC_0x01_0x14_OServiceSetPrivacyFlags struct {
 	PrivacyFlags uint32
 	PrivacyFlags uint32
 }
 }