summaryrefslogtreecommitdiffstats
path: root/katzenpost/cmd/yamn-dispatcher/main.go
blob: 8f10034dbc2a2051070594b581cbbd3cbfb13686 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
package main

import (
	"errors"
	"flag"
	"fmt"
	"os"
	"path/filepath"
	"strings"
	"time"

	"git.virebent.art/virebent/yamnweb/katzenpost/dispatcher"
	"github.com/katzenpost/katzenpost/core/log"
	"github.com/katzenpost/katzenpost/server/cborplugin"
)

const capability = "yamn-dispatch-v1"

type plugin struct {
	write          func(cborplugin.Command)
	reassembler    *dispatcher.Reassembler
	allowedDomains map[string]struct{}
}

func (p *plugin) OnCommand(command cborplugin.Command) error {
	request, ok := command.(*cborplugin.Request)
	if !ok {
		return errors.New("unexpected plugin command")
	}
	frame, err := dispatcher.DecodeFrame(request.Payload)
	if err != nil {
		return err
	}
	encoded, complete, err := p.reassembler.Add(frame)
	if err != nil || !complete {
		return err
	}
	envelope, err := dispatcher.DecodeEnvelope(encoded)
	if err != nil {
		return err
	}
	if err := envelope.Validate(p.allowedDomains); err != nil {
		return err
	}
	// PoC sink: successful validation intentionally has no external side effect.
	return nil
}

func (p *plugin) RegisterConsumer(server *cborplugin.Server) {
	p.write = server.Write
}

func parseDomains(value string) (map[string]struct{}, error) {
	domains := make(map[string]struct{})
	for _, domain := range strings.Split(value, ",") {
		domain = strings.ToLower(strings.TrimSpace(domain))
		if domain == "" || strings.ContainsAny(domain, "@/\\: ") {
			return nil, errors.New("invalid allowed domain")
		}
		domains[domain] = struct{}{}
	}
	if len(domains) == 0 {
		return nil, errors.New("at least one allowed domain is required")
	}
	return domains, nil
}

func run() error {
	var allowed string
	var logDir string
	var logLevel string
	flag.StringVar(&allowed, "allowed-domains", "remailer.example", "comma-separated entry remailer domains")
	flag.StringVar(&logDir, "log-dir", "/tmp", "operational log directory")
	flag.StringVar(&logLevel, "log-level", "NOTICE", "Katzenpost log level")
	flag.Parse()

	domains, err := parseDomains(allowed)
	if err != nil {
		return err
	}
	info, err := os.Stat(logDir)
	if err != nil || !info.IsDir() {
		return errors.New("log directory is unavailable")
	}
	backend, err := log.New(filepath.Join(logDir, "yamn-dispatch.log"), logLevel, false)
	if err != nil {
		return fmt.Errorf("initialize logging: %w", err)
	}
	logger := backend.GetLogger("yamn_dispatch")

	socketDir, err := os.MkdirTemp("", "yamn-dispatch-")
	if err != nil {
		return fmt.Errorf("create socket directory: %w", err)
	}
	defer os.RemoveAll(socketDir)
	socketPath := filepath.Join(socketDir, "plugin.sock")
	service := &plugin{
		reassembler:    dispatcher.NewReassembler(5 * time.Minute),
		allowedDomains: domains,
	}
	server := cborplugin.NewServer(logger, socketPath, new(cborplugin.RequestFactory), service)
	if _, err := fmt.Fprintln(os.Stdout, socketPath); err != nil {
		return fmt.Errorf("publish socket path: %w", err)
	}
	server.Accept()
	server.Wait()
	return nil
}

func main() {
	if err := run(); err != nil {
		fmt.Fprintln(os.Stderr, "yamn-dispatcher failed")
		os.Exit(1)
	}
}