Files
go-sip/internal/asterisk/loader.go
T

285 lines
8.8 KiB
Go

package asterisk
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
"os/exec"
"path/filepath"
"regexp"
"sort"
"strings"
"time"
"git.ipao.vip/rogee/go-sip/internal/configread"
)
// Loader writes only the Agent-owned endpoint include. A reviewed static
// pjsip.conf owns the transport, whose reload can interrupt live calls.
type Loader struct {
ConfigDir string
Asterisk string
LibraryDir string
}
type loadedTrunk struct {
ID string `json:"id"`
Host string `json:"host"`
Port int `json:"port"`
}
type appliedState struct {
Revision int64 `json:"revision"`
Hash string `json:"hash"`
SnapshotHash string `json:"snapshot_hash"`
Trunks []loadedTrunk `json:"trunks"`
}
var endpointLine = regexp.MustCompile(`(?m)^\s*Endpoint:\s+(\S+)`)
func (l Loader) paths() (string, string, error) {
if l.ConfigDir == "" || l.Asterisk == "" || l.LibraryDir == "" {
return "", "", errors.New("native Asterisk executable, library and config paths are required")
}
base, err := os.ReadFile(filepath.Join(l.ConfigDir, "pjsip.conf"))
if err != nil {
return "", "", fmt.Errorf("read management-owned PJSIP base: %w", err)
}
if !bytes.Contains(base, []byte("#tryinclude go-sip-managed.conf")) || !bytes.Contains(base, []byte("[go-sip-udp]")) {
return "", "", errors.New("PJSIP base is missing the approved static transport and managed include")
}
return filepath.Join(l.ConfigDir, "go-sip-managed.conf"), filepath.Join(l.ConfigDir, "go-sip-applied.json"), nil
}
func (l Loader) command(ctx context.Context, name string, args ...string) ([]byte, error) {
cmd := exec.CommandContext(ctx, name, args...)
cmd.Env = append(os.Environ(), "LD_LIBRARY_PATH="+l.LibraryDir)
out, err := cmd.CombinedOutput()
if err != nil {
return nil, fmt.Errorf("native Asterisk command %q failed: %w (output_sha256=%x)", filepath.Base(name), err, sha256.Sum256(out))
}
return out, nil
}
func (l Loader) cli(ctx context.Context, command string) ([]byte, error) {
return l.command(ctx, l.Asterisk, "-C", filepath.Join(l.ConfigDir, "asterisk.conf"), "-rx", command)
}
func (l Loader) observe(ctx context.Context, state appliedState) error {
if state.Revision < 1 || state.SnapshotHash == "" || len(state.Trunks) == 0 {
return errors.New("no verified active SIP trunks")
}
transport, err := l.cli(ctx, "pjsip show transports")
if err != nil {
return fmt.Errorf("inspect native UDP transport: %w", err)
}
approvedTransport := false
for _, line := range strings.Split(string(transport), "\n") {
fields := strings.Fields(line)
if len(fields) >= 6 && fields[0] == "Transport:" && fields[1] == "go-sip-udp" && fields[2] == "udp" && fields[len(fields)-1] == "0.0.0.0:5060" {
approvedTransport = true
}
}
if !approvedTransport {
return errors.New("approved native UDP transport/address is not loaded")
}
all, err := l.cli(ctx, "pjsip show endpoints")
if err != nil {
return err
}
expected := make(map[string]loadedTrunk, len(state.Trunks))
for _, t := range state.Trunks {
expected[t.ID] = t
}
seen := make(map[string]bool)
for _, match := range endpointLine.FindAllSubmatch(all, -1) {
id := string(match[1])
if !trunkName.MatchString(id) { // skip Asterisk's <Endpoint/CID> heading
continue
}
if _, ok := expected[id]; !ok {
return fmt.Errorf("unapproved Asterisk endpoint %q is loaded", id)
}
seen[id] = true
}
if len(seen) != len(expected) {
return errors.New("approved SIP endpoint set is not loaded")
}
for _, t := range state.Trunks {
endpoint, err := l.cli(ctx, "pjsip show endpoint "+t.ID)
if err != nil || !bytes.Contains(endpoint, []byte(t.ID+"-aor")) || !bytes.Contains(endpoint, []byte("alaw")) || !bytes.Contains(endpoint, []byte("go-sip-udp")) || !bytes.Contains(endpoint, []byte("go-sip-no-inbound")) {
return fmt.Errorf("SIP endpoint %q did not load its approved AOR/codec", t.ID)
}
aor, err := l.cli(ctx, "pjsip show aor "+t.ID+"-aor")
if err != nil || !bytes.Contains(aor, []byte(fmt.Sprintf("sip:%s:%d", t.Host, t.Port))) {
return fmt.Errorf("SIP AOR %q did not load its approved contact", t.ID)
}
}
return nil
}
// LoadedSIP never trusts a persisted revision alone: it also checks the file
// hash and live PJSIP objects after every Agent boot or Dispatcher query.
func (l Loader) LoadedSIP(ctx context.Context) (map[string]int64, error) {
config, statePath, err := l.paths()
if err != nil {
return nil, err
}
raw, err := os.ReadFile(statePath)
if err != nil {
return nil, fmt.Errorf("read last verified SIP revision: %w", err)
}
var state appliedState
if err := json.Unmarshal(raw, &state); err != nil {
return nil, fmt.Errorf("decode last verified SIP revision: %w", err)
}
data, err := os.ReadFile(config)
if err != nil {
return nil, err
}
hash := sha256.Sum256(data)
if hex.EncodeToString(hash[:]) != state.Hash {
return nil, errors.New("SIP config no longer matches last verified revision")
}
if err := l.observe(ctx, state); err != nil {
return nil, err
}
loaded := make(map[string]int64, len(state.Trunks))
for _, t := range state.Trunks {
loaded[t.ID] = state.Revision
}
return loaded, nil
}
// Apply changes no transport, routes or existing call. It records a revision
// only after the native module reload and readback succeed.
func (l Loader) Apply(ctx context.Context, approvedJSON []byte) (map[string]int64, error) {
config, statePath, err := l.paths()
if err != nil {
return nil, err
}
var sip configread.SIP
if err := json.Unmarshal(approvedJSON, &sip); err != nil {
return nil, err
}
// The same renderer is used by RPC validation and disk application.
text, err := Render(sip)
if err != nil {
return nil, err
}
var trunks []struct {
ID string `json:"trunk_id"`
Host string `json:"server_host"`
Port int `json:"server_port"`
Enabled bool `json:"enabled"`
}
if err := json.Unmarshal(sip.Trunks, &trunks); err != nil {
return nil, err
}
canonical, err := json.Marshal(sip)
if err != nil {
return nil, err
}
state := appliedState{Revision: sip.Revision, Hash: fmt.Sprintf("%x", sha256.Sum256([]byte(text))), SnapshotHash: fmt.Sprintf("%x", sha256.Sum256(canonical))}
for _, t := range trunks {
if t.Enabled {
state.Trunks = append(state.Trunks, loadedTrunk{ID: t.ID, Host: t.Host, Port: t.Port})
}
}
sort.Slice(state.Trunks, func(i, j int) bool { return state.Trunks[i].ID < state.Trunks[j].ID })
recoverExisting := false
if existing, err := os.ReadFile(statePath); err == nil {
var previous appliedState
if err := json.Unmarshal(existing, &previous); err != nil {
return nil, err
}
if sip.Revision < previous.Revision || (sip.Revision == previous.Revision && (state.Hash != previous.Hash || state.SnapshotHash != previous.SnapshotHash)) {
return nil, errors.New("approved SIP revision regressed or changed content")
}
if sip.Revision == previous.Revision {
return l.LoadedSIP(ctx)
}
} else if !errors.Is(err, os.ErrNotExist) {
return nil, err
} else if current, err := os.ReadFile(config); err == nil {
if sha256.Sum256(current) != sha256.Sum256([]byte(text)) {
return nil, errors.New("unverified managed SIP config does not match the approved snapshot")
}
recoverExisting = true
} else if !errors.Is(err, os.ErrNotExist) {
return nil, err
}
if !recoverExisting {
if err := writeAtomic(config, []byte(text)); err != nil {
return nil, err
}
if _, err := l.command(ctx, "systemctl", "--user", "reload", "go-sip-asterisk.service"); err != nil {
return nil, err
}
}
if err := l.waitObserved(ctx, state); err != nil {
return nil, err
}
encoded, err := json.Marshal(state)
if err != nil {
return nil, err
}
if err := writeAtomic(statePath, encoded); err != nil {
return nil, err
}
return l.LoadedSIP(ctx)
}
func (l Loader) waitObserved(ctx context.Context, state appliedState) error {
deadline := time.NewTimer(15 * time.Second)
defer deadline.Stop()
var last error
for attempts := 1; ; attempts++ {
if last = l.observe(ctx, state); last == nil {
return nil
}
select {
case <-ctx.Done():
return fmt.Errorf("native SIP readback cancelled after %d attempts: %w", attempts, ctx.Err())
case <-deadline.C:
return fmt.Errorf("native SIP readback timed out after %d attempts: %w", attempts, last)
case <-time.After(250 * time.Millisecond):
}
}
}
func writeAtomic(path string, data []byte) error {
file, err := os.CreateTemp(filepath.Dir(path), ".go-sip-stage-*")
if err != nil {
return err
}
defer os.Remove(file.Name())
defer file.Close()
if err := file.Chmod(0600); err != nil {
return err
}
if _, err := file.Write(data); err != nil {
return err
}
if err := file.Sync(); err != nil {
return err
}
if err := file.Close(); err != nil {
return err
}
if err := os.Rename(file.Name(), path); err != nil {
return err
}
dir, err := os.Open(filepath.Dir(path))
if err != nil {
return err
}
defer dir.Close()
return dir.Sync()
}