summaryrefslogtreecommitdiff
path: root/internal/server/server.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/server/server.go')
-rw-r--r--internal/server/server.go97
1 files changed, 82 insertions, 15 deletions
diff --git a/internal/server/server.go b/internal/server/server.go
index 32c1205..cdb9313 100644
--- a/internal/server/server.go
+++ b/internal/server/server.go
@@ -83,11 +83,15 @@ func (s *server) AddSession(session *session.Session, userId snowflake.ID, pubKe
if sess, ok := s.sessions[session.ID()]; ok {
EvictSession(sess) // last connection wins
slog.Info("closed due to new connection from another location",
- ctxkeys.IpAddr, sess.Addr(), ctxkeys.UserID, sess.ID(),
- ctxkeys.EvictedBy, session.Addr())
+ ctxkeys.IpAddr.String(), sess.Addr(),
+ ctxkeys.UserID.String(), sess.ID(),
+ ctxkeys.EvictedBy.String(), session.Addr(),
+ )
slog.Info("this session evicted another session",
- ctxkeys.IpAddr, session.Addr(), ctxkeys.UserID, session.ID(),
- ctxkeys.Evicted, sess.Addr())
+ ctxkeys.IpAddr.String(), session.Addr(),
+ ctxkeys.UserID.String(), session.ID(),
+ ctxkeys.Evicted.String(), sess.Addr(),
+ )
}
s.sessions[session.ID()] = session
@@ -98,8 +102,7 @@ func EvictSession(sess *session.Session) {
ctx, cancel := context.WithTimeout(context.Background(), timeout)
defer cancel()
payload := &packet.Error{
- Error: "new connection from another location, closing this one",
- PktType: packet.PacketError,
+ Error: "new connection from another location, closing this one",
}
pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload))
sess.Write(ctx, pkt)
@@ -206,6 +209,7 @@ func (server *server) handleConnection(conn net.Conn) {
// Writer
go func() {
defer close(done)
+ defer conn.Close() // To unblock reader
writeQueue := sess.Read()
for packet := range writeQueue {
@@ -242,6 +246,10 @@ func (server *server) handleConnection(conn net.Conn) {
}
}()
+ // NOTE: IMPROTANT LEGAL STUFF
+ // Sending this first thing, before client sends us any data
+ sendTosInfo(ctx, sess)
+
// Reader
buffer := make([]byte, 512)
for {
@@ -328,11 +336,60 @@ func processPacket(ctx context.Context, sess *session.Session, pkt packet.Packet
}
func processRequest(ctx context.Context, sess *session.Session, request packet.Payload) packet.Payload {
- log.Println(sess.Addr(), "processing", request.Type(), "request:", request)
+ slog.InfoContext(ctx, "processing request",
+ ctxkeys.RequestType.String(),
+ request.Type(), ctxkeys.Request.String(), request,
+ )
+
+ if !sess.IsTosAccepted() {
+ if acceptTos, ok := request.(*packet.AcceptTos); ok && acceptTos.IAgreeToTheTermsOfServiceAndPrivacyPolicy {
+ sess.ReceivedTosAcceptance()
+ slog.InfoContext(ctx, "terms of service accepted, continuing...")
+ return &api.ErrSuccess
+ }
+
+ slog.InfoContext(ctx, "refused terms of service, refusing service...")
+ sess.Close() // Refuse to receive any more requests
+ return &packet.Error{Error: "Terms of Service not accepted, refusing service"}
+ }
+
+ assert.Assert(sess.IsTosAccepted(), "justified paranoia") // Just in case
// TODO: add a way to measure the time each request/response took and log it
// Potentially even separate time for code vs DB operations
+
var response packet.Payload
+
+ if sess.IsAuthenticated() {
+ // Authentication only requests, others will be handled without auth even if authenticated
+ authCtx := ctxkeys.WithValue(ctx, ctxkeys.UserID, sess.ID())
+ response = processAuthenticatedRequests(authCtx, sess, request)
+ }
+
+ if response != nil {
+ return response
+ }
+
+ switch request := request.(type) {
+ // TODO:
+ // case *packet.GetNonce:
+ // response = timeout(5*time.Millisecond, api.SetUserData, ctx, sess, request)
+
+ // TODO:
+ // case *packet.Authenticate:
+ // response = timeout(5*time.Millisecond, api.SetUserData, ctx, sess, request)
+
+ default:
+ _ = request // FIXME: remove
+ response = &packet.Error{Error: "use of disallowed packet type for request"}
+ }
+
+ return response
+}
+
+func processAuthenticatedRequests(ctx context.Context, sess *session.Session, request packet.Payload) packet.Payload {
+ var response packet.Payload
+
switch request := request.(type) {
case *packet.SetUserData:
@@ -383,13 +440,7 @@ func processRequest(ctx context.Context, sess *session.Session, request packet.P
response = timeout(10*time.Millisecond, api.GetUsers, ctx, sess, request)
default:
- response = &packet.Error{Error: "use of disallowed packet type for request"}
- }
-
- // TODO: Isn't this weird? shouldn't it always be packet.ErrorPacket type?
- // rather than copying the pkt type of the request?
- if response, ok := response.(*packet.Error); ok {
- response.PktType = request.Type()
+ response = nil
}
return response
@@ -422,7 +473,23 @@ func timeout[T packet.Payload](
}
}
-func (server *server) sendInitialPackets(ctx context.Context, sess *session.Session) bool {
+func sendTosInfo(ctx context.Context, sess *session.Session) bool {
+ // FIXME: ===== REPLACE WITH ACTUAL TOS =====
+ // TODO: Get TOS and privacy policy
+ tos := "WIP Terms of Service"
+ privacy := "WIP Privacy Policy"
+ date := "2025-07-03"
+
+ payload := &packet.TosInfo{
+ Tos: tos,
+ PrivacyPolicy: privacy,
+ Date: date,
+ }
+ pkt := packet.NewPacket(packet.NewMsgPackEncoder(payload))
+ return sess.Write(ctx, pkt)
+}
+
+func (server *server) sendInitialAuthPackets(ctx context.Context, sess *session.Session) bool {
payload := api.GetUserData(ctx, sess, &packet.GetUserData{})
if payload == &api.ErrInternalError {
return false