diff options
Diffstat (limited to 'internal')
| -rw-r--r-- | internal/server/ctxkeys/ctxkeys.go | 51 | ||||
| -rw-r--r-- | internal/server/server.go | 7 |
2 files changed, 58 insertions, 0 deletions
diff --git a/internal/server/ctxkeys/ctxkeys.go b/internal/server/ctxkeys/ctxkeys.go new file mode 100644 index 0000000..c0f688d --- /dev/null +++ b/internal/server/ctxkeys/ctxkeys.go @@ -0,0 +1,51 @@ +package ctxkeys + +import ( + "context" + "log/slog" + + "github.com/kyren223/eko/pkg/assert" +) + +type key int + +const ( + UserID key = iota + IpAddr + KeyMax +) + +var keyNames = map[key]string{ + UserID: "user_id", + IpAddr: "ip_addr", +} + +func Init() { + assert.Assert(len(keyNames) == int(KeyMax), "Keys in keyNames mismatch amount of keys", "len(keyNames)", len(keyNames), "KeyMax", int(KeyMax)) +} + +func WithValue(ctx context.Context, k key, v any) context.Context { + return context.WithValue(ctx, k, v) +} + +func Value(ctx context.Context, k key) any { + return ctx.Value(k) +} + +type ContextHandler struct { + slog.Handler +} + +func WrapLogHandler(baseHandler slog.Handler) *ContextHandler { + return &ContextHandler{Handler: baseHandler} +} + +func (h *ContextHandler) Handle(ctx context.Context, r slog.Record) error { + for k := key(0); k < KeyMax; k++ { + if v := Value(ctx, k); v != nil { + r.AddAttrs(slog.Any(keyNames[k], v)) + } + } + + return h.Handler.Handle(ctx, r) +} diff --git a/internal/server/server.go b/internal/server/server.go index 8067a19..8accb14 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -19,6 +19,7 @@ import ( "github.com/kyren223/eko/certs" "github.com/kyren223/eko/internal/packet" "github.com/kyren223/eko/internal/server/api" + "github.com/kyren223/eko/internal/server/ctxkeys" "github.com/kyren223/eko/internal/server/session" "github.com/kyren223/eko/pkg/assert" "github.com/kyren223/eko/pkg/snowflake" @@ -151,6 +152,7 @@ func (server *server) handleConnection(conn net.Conn) { log.Println(addr, "accepted") initialCtx, initialCancel := context.WithTimeout(server.ctx, 5*time.Second) + initialCtx = context.WithValue(initialCtx, ctxkeys.IpAddr, addr) deadline, _ := initialCtx.Deadline() err := conn.SetDeadline(deadline) assert.NoError(err, "setting read deadline should not error") @@ -175,8 +177,13 @@ func (server *server) handleConnection(conn net.Conn) { log.Println(addr, "disconnected") return } + ctx, cancel := context.WithCancel(server.ctx) defer cancel() + + ctx = context.WithValue(ctx, ctxkeys.UserID, user.ID) + ctx = context.WithValue(ctx, ctxkeys.IpAddr, addr) + sess := session.NewSession(server, addr, cancel, user.ID, pubKey) server.AddSession(sess) framer := packet.NewFramer() |
