19 Commits
Author SHA1 Message Date
MengMengCode 4df0ae0c7d feat: expand device networking and management 2026-08-10 03:27:43 +08:00
MengMengCode d8828ff26a fix: expose IMS call failures and allow unlimited calls 2026-08-09 21:52:25 +08:00
MengMengCode 337aa3c0ab fix: route calls only through ready IMS sessions 2026-08-09 21:31:24 +08:00
MengMengCode 97ca84bbfc fix: redact Telegram tokens from errors 2026-08-09 21:14:07 +08:00
MengMengCode cc477571ac feat: add Telegram AT and USSD commands 2026-08-09 21:06:35 +08:00
MengMengCode c19156e46a fix: redirect expired sessions to login 2026-08-09 20:47:53 +08:00
MengMengCode 0e68dc6893 fix: revoke sessions after updates 2026-08-09 20:23:00 +08:00
MengMengCode 1fc6ea9b6c fix: make web updates restart cleanly 2026-08-09 20:13:50 +08:00
MengMengCode 85f8790e1e fix: allow verified in-place updates 2026-08-09 20:09:44 +08:00
MengMengCode 3d117749b2 build: speed up multi-arch Docker builds 2026-08-09 19:59:05 +08:00
MengMengCode 3604319faa feat: update SMS handling to include readyForClose channel for better session management 2026-08-09 19:35:00 +08:00
MengMengCode 9fc3f1c5b8 feat: add Telegram API URL handling and version comparison utilities
- Implemented `telegramAPIURL` and `validateTelegramAPIURL` functions for constructing and validating Telegram API URLs.
- Added semantic versioning utilities in `version.go` to compare versions and validate semantic version formats.
- Created tests for version comparison logic in `version_test.go`.
- Introduced IMS call handling in `call_runtime.go`, including methods for dialing, answering, and hanging up calls.
- Added tests for incoming call handling and validation in `call_runtime_test.go`.
- Developed a new `PluginsCard` component for managing plugins via URL or file upload in the web interface.
- Implemented plugin management functions in `extensions.ts` and created an `ExtensionPage` for displaying plugin contributions.
2026-08-09 19:28:12 +08:00
MengMengCode 93b0cf718c Merge branch 'master' of https://github.com/MengMengCode/VoCat 2026-08-09 16:59:32 +08:00
MengMengCode 3939b061af Enhance session relaying to handle delayed IKE packets and implement an automatic retry mechanism 2026-08-09 16:59:28 +08:00
Meng MengandGitHub 03c8a2ceae Update README.md 2026-08-09 16:38:56 +08:00
Meng MengandGitHub d7a9fc9774 Add group and channel information
Added group and channel links to README.
2026-08-09 16:38:05 +08:00
Meng MengandGitHub 147721f237 Update thanks section in README.md 2026-08-09 16:25:14 +08:00
Meng MengandGitHub e24be6ef29 Add 'Thanks' section to README
Added a 'Thanks' section with community links.
2026-08-09 16:21:18 +08:00
MengMengCode 6f2b7bf395 Initial 2026-08-09 15:53:35 +08:00
153 changed files with 11970 additions and 605 deletions
+6
View File
@@ -0,0 +1,6 @@
# Copy this file to .env and fill in real values before `docker compose up -d`.
# .env is gitignored; .env.example is tracked as a template.
# Admin password for the web UI. REQUIRED — the server refuses to start safely
# without it once exposed. Pick a strong password.
VOCAT_ADMIN_PASSWORD=change-me-to-a-strong-password
+4
View File
@@ -72,6 +72,10 @@ jobs:
goarch: arm64
goarm: ""
filename: vocat-linux-arm64
- target: linux-aarch64
goarch: arm64
goarm: ""
filename: vocat-linux-aarch64
- target: linux-armv7
goarch: arm
goarm: "7"
+6
View File
@@ -15,6 +15,9 @@
vc.jar
*.cookies
*.session
.env
.env.*
!.env.example
# ---- Frontend build products ----
web/dist/
@@ -33,9 +36,12 @@ __pycache__/
# committed. The whole build directory and the root release.py are ignored.
build/*.py
/extension/
# ---- Docs / scratch ----
*.md
!README.md
!docs/**/*.md
*.txt
build/lists/
+13 -5
View File
@@ -1,20 +1,25 @@
# syntax=docker/dockerfile:1.7
# ---- Stage 1: build the web frontend ----
FROM node:20-alpine AS web-builder
# Build toolchains run natively on the BuildKit host. Without BUILDPLATFORM,
# the arm64 branch executes npm and the Go compiler through QEMU, which is much
# slower and makes npm ci appear to hang despite producing no progress output.
# ---- Stage 1: build the web frontend once on the native builder ----
FROM --platform=$BUILDPLATFORM node:20-alpine AS web-builder
WORKDIR /web
COPY web/package.json web/package-lock.json* ./
RUN npm ci
COPY web/ ./
RUN npm run build
# ---- Stage 2: build the Go binary ----
FROM golang:1.25-alpine AS go-builder
# ---- Stage 2: cross-compile the Go binary on the native builder ----
FROM --platform=$BUILDPLATFORM golang:1.25-alpine AS go-builder
RUN apk add --no-cache git
WORKDIR /src
ARG VERSION=0.1.0-dev
ARG BUILD_TIME=""
ARG TARGETOS
ARG TARGETARCH
COPY go.mod go.sum ./
RUN go mod download
@@ -23,7 +28,7 @@ COPY . .
# Overlay the freshly built frontend so go:embed web/dist picks it up.
COPY --from=web-builder /web/dist ./web/dist
RUN CGO_ENABLED=0 GOOS=linux go build \
RUN CGO_ENABLED=0 GOOS=${TARGETOS:-linux} GOARCH=${TARGETARCH} go build \
-trimpath \
-ldflags "-s -w -X vocat/internal/buildinfo.Version=${VERSION} -X vocat/internal/buildinfo.BuildTime=${BUILD_TIME}" \
-o /out/vocat \
@@ -40,6 +45,9 @@ RUN mkdir -p /opt/vocat/bin /opt/vocat/data && \
COPY --from=go-builder /out/vocat /opt/vocat/bin/vocat
# Symlink into /usr/local/bin so `docker exec <ctr> vocat ...` finds it via $PATH.
RUN ln -s /opt/vocat/bin/vocat /usr/local/bin/vocat
USER vocat
VOLUME ["/opt/vocat/data"]
EXPOSE 7575
+25 -9
View File
@@ -14,7 +14,7 @@
</p>
<p align="center">
<img alt="Linux" src="https://img.shields.io/badge/Linux-amd64_%7C_386_%7C_arm64_%7C_armv7-FCC624?style=flat-square&logo=linux&logoColor=111111">
<img alt="Linux" src="https://img.shields.io/badge/Linux-amd64_%7C_386_%7C_arm64_%7C_aarch64_%7C_armv7-FCC624?style=flat-square&logo=linux&logoColor=111111">
<img alt="Docker" src="https://img.shields.io/badge/Docker-Multi--Arch-2496ED?style=flat-square&logo=docker&logoColor=white">
<img alt="WiFi Calling" src="https://img.shields.io/badge/WiFi_Calling-IMS_SMS-7B1FA2?style=flat-square">
<img alt="eSIM" src="https://img.shields.io/badge/eSIM-LPA_%2F_eUICC-009688?style=flat-square">
@@ -22,6 +22,8 @@
<img alt="GitHub Actions" src="https://img.shields.io/badge/GitHub_Actions-Release-2088FF?style=flat-square&logo=githubactions&logoColor=white">
</p>
**English** | [简体中文](docs/README.zh-CN.md)
Vocat is an open-source web control panel and engineering toolkit for Quectel EC20/EC25-class cellular modems. It combines modem discovery, live radio status, AT and USSD terminals, SMS, WiFi Calling, eSIM management, network selection, proxy routing, notifications, audit logs, and release automation in one self-contained service.
The backend is written in Go, the interface is built with React and TypeScript, and the production frontend is embedded into the Go binary. A single executable contains the web application and uses SQLite for persistent state.
@@ -64,23 +66,23 @@ Available features depend on the module firmware, USB composition, SIM/eSIM capa
### One-click Linux installation
```bash
curl -fsSL https://raw.githubusercontent.com/MengMengCode/VoCat/main/scripts/install.sh | sudo bash
curl -fsSL https://raw.githubusercontent.com/MengMengCode/VoCat/master/scripts/install.sh | sudo bash
```
Install a specific version:
```bash
curl -fsSL https://raw.githubusercontent.com/MengMengCode/VoCat/main/scripts/install.sh -o install.sh
sudo bash install.sh 0.2.0
curl -fsSL https://raw.githubusercontent.com/MengMengCode/VoCat/master/scripts/install.sh -o install.sh
sudo bash install.sh 0.0.2
```
The installer:
- detects `amd64`, `386`, `arm64`, or `armv7`;
- detects `amd64`, `386`, `arm64`, `aarch64`, or `armv7`;
- downloads the matching GitHub Release binary;
- verifies it against `SHA256SUMS`;
- installs Vocat under `/opt/vocat`;
- creates a dedicated system user and systemd service;
- creates a hardened systemd service with the hardware and network access required by Vocat;
- stores runtime configuration in `/etc/vocat/env`;
- generates a random initial administrator password on first installation.
@@ -99,6 +101,7 @@ Download the matching binary and `SHA256SUMS` from GitHub Releases:
| Linux x86-64 | `vocat-linux-amd64` |
| Linux x86 32-bit | `vocat-linux-386` |
| Linux ARM64 | `vocat-linux-arm64` |
| Linux AArch64 | `vocat-linux-aarch64` |
| Linux ARMv7 | `vocat-linux-armv7` |
Verify and install it:
@@ -107,9 +110,17 @@ Verify and install it:
sha256sum -c SHA256SUMS --ignore-missing
sudo install -d -m 0755 /opt/vocat/bin /opt/vocat/data
sudo install -m 0755 vocat-linux-amd64 /opt/vocat/bin/vocat
sudo /opt/vocat/bin/vocat
sudo env \
VOCAT_DATABASE_PATH=/opt/vocat/data/vocat.db \
VOCAT_ADMIN_PASSWORD=change-this-password \
/opt/vocat/bin/vocat serve
```
This manual command runs Vocat in the foreground. Use `vocat serve` so the
process starts the server directly; running `vocat` without arguments as root
on a TTY opens the interactive management menu instead. Use the one-click
installer when a managed systemd service and automatic restart are required.
### Docker
For a Linux host that must discover every attached supported Quectel modem and
@@ -161,7 +172,7 @@ Vocat reads an optional JSON configuration file from `VOCAT_CONFIG`, then applie
| `VOCAT_SECURE_COOKIES` | `false` | Marks session cookies as secure when HTTPS is used. |
| `VOCAT_SHUTDOWN_TIMEOUT` | `10s` | Graceful shutdown timeout. |
| `VOCAT_MAX_REQUEST_BODY_BYTES` | `1048576` | Maximum API request body size. |
| `VOCAT_REPO` | empty | GitHub repository used by the self-updater, in `owner/name` form. |
| `VOCAT_REPO` | `MengMengCode/VoCat` | Trusted GitHub repository used by the self-updater, in `owner/name` form. |
| `GITHUB_TOKEN` | empty | Optional GitHub token for private repositories or higher API limits. |
Do not store Telegram tokens, SMTP passwords, webhook secrets, SIM credentials, or other private data in the repository. Configure them through the application settings or protected environment files.
@@ -249,7 +260,7 @@ go build -trimpath -ldflags "-s -w" -o vocat ./cmd/vocat
Pushing a version tag starts two GitHub Actions workflows:
- `release-binaries` builds and publishes `amd64`, `386`, `arm64`, and `armv7` binaries plus `SHA256SUMS`.
- `release-binaries` builds and publishes `amd64`, `386`, `arm64`, `aarch64`, and `armv7` binaries plus `SHA256SUMS`.
- `docker` builds and publishes a multi-architecture image to GitHub Container Registry.
```bash
@@ -289,6 +300,11 @@ go test ./...
cd web && npm run build
```
## Thanks
- [Nodeseek.com](https://www.nodeseek.com) — A community dedicated to servers
- [Linux.do](https://linux.do) — An inspiring tech community
- [iniwex5](https://github.com/iniwex5) - Style and Functionality Guidelines
## License
See [LICENSE](LICENSE).
+12 -6
View File
@@ -15,23 +15,29 @@ func printUsage(w io.Writer) {
fmt.Fprintf(w, `vocat %s
Usage:
vocat Run the vocat server (default; same as no arguments).
vocat No arguments: interactive management menu when run as
root on a TTY, otherwise the server. systemd (non-TTY)
starts the server unchanged.
vocat serve Run the server in the foreground (use from a TTY when
vocat without arguments would enter the menu).
vocat version Print the build version and exit.
vocat update Check GitHub for a newer release and self-update.
Flags:
--check Only report whether an update is available.
--repo owner/name GitHub repository (default: $VOCAT_REPO).
--repo owner/name GitHub repository (default: $VOCAT_REPO or MengMengCode/VoCat).
--target path Binary to replace (default: running exe).
--force Reinstall even at the same version.
Environment:
VOCAT_REPO Fallback for --repo.
GITHUB_TOKEN Optional bearer token for private repos
or higher rate limits.
vocat menu Interactive lifecycle menu (run as root on the host):
change password, restart service, uninstall.
vocat menu Interactive lifecycle menu (root on the host):
toggle language, change password, restart, update,
uninstall.
vocat help Show this help message.
When run without a subcommand, vocat starts the HTTP server using
VOCAT_* environment variables or $VOCAT_CONFIG for configuration.
When run without a subcommand on a non-TTY (e.g. systemd), vocat starts the
HTTP server using VOCAT_* environment variables or $VOCAT_CONFIG for
configuration.
`, buildinfo.Version)
}
+100
View File
@@ -0,0 +1,100 @@
package main
import (
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
"strings"
"time"
"vocat/internal/config"
"vocat/internal/developer"
"vocat/internal/store"
)
// developerEnabledSettingKey is the app_settings key that gates the entire
// plugin/extension system. When absent the developer mode defaults to off, so
// a fresh install exposes no plugin surface until an operator explicitly turns
// it on with `vocat develop on` and restarts the service.
const developerEnabledSettingKey = developer.EnabledSettingKey
// runDevelop handles the hidden `vocat develop on|off` subcommand. It is
// intentionally excluded from printUsage and the interactive menu: the plugin
// system is an opt-in developer surface, and the toggle must be typed in full
// to activate it. The flag is persisted to app_settings and takes effect on
// the next server start (run() reads it before creating the plugin manager).
func runDevelop(args []string, logger *slog.Logger) error {
if len(args) == 0 {
return errors.New(`usage: vocat develop <on|off>`)
}
enabled, ok := parseDevelopArg(args[0])
if !ok {
return fmt.Errorf(`vocat develop: invalid argument %q (expected "on" or "off")`, args[0])
}
// Match the menu's env resolution. An operator runs `vocat develop` on the
// host where the shell has not sourced /etc/vocat/env (a systemd
// EnvironmentFile, not a shell rc) and VOCAT_DATABASE_PATH is unset, so
// config.Load() would otherwise resolve a CWD-relative ./data/vocat.db — a
// different database than /opt/vocat/data/vocat.db the service reads. The
// flag would then be written to a DB the service never opens, silently.
loadMenuEnv()
cfg, err := config.Load()
if err != nil {
return fmt.Errorf("load configuration: %w", err)
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
database, err := store.Open(ctx, cfg.DatabasePath)
if err != nil {
return fmt.Errorf("open database: %w", err)
}
defer database.Close()
payload, err := json.Marshal(map[string]bool{"enabled": enabled})
if err != nil {
return fmt.Errorf("encode developer flag: %w", err)
}
if err := database.UpsertAppSetting(ctx, store.AppSetting{
Key: developerEnabledSettingKey,
Value: payload,
}); err != nil {
return fmt.Errorf("persist developer flag: %w", err)
}
if !enabled {
if err := developer.ResetExperimental(ctx, database); err != nil {
return fmt.Errorf("reset developer settings: %w", err)
}
}
if enabled {
fmt.Printf("开发者模式已开启。重启 vocat 服务后插件功能生效。\n数据库:%s\n", cfg.DatabasePath)
fmt.Printf("Developer mode enabled. Restart the vocat service for plugins to take effect.\nDatabase: %s\n", cfg.DatabasePath)
} else {
fmt.Printf("开发者模式已关闭。重启 vocat 服务后插件功能将停用。\n数据库:%s\n", cfg.DatabasePath)
fmt.Printf("Developer mode disabled. Restart the vocat service to deactivate plugins.\nDatabase: %s\n", cfg.DatabasePath)
}
return nil
}
func parseDevelopArg(arg string) (bool, bool) {
switch strings.ToLower(strings.TrimSpace(arg)) {
case "on", "1", "true", "yes":
return true, true
case "off", "0", "false", "no":
return false, true
default:
return false, false
}
}
// isDeveloperEnabled reads the persisted developer-mode flag. A missing record
// or an unparseable value resolves to false — the system defaults closed, so
// any read failure keeps plugins off rather than exposing them by accident.
func isDeveloperEnabled(ctx context.Context, database *store.Store) bool {
return developer.Enabled(ctx, database)
}
+304 -18
View File
@@ -2,20 +2,29 @@ package main
import (
"context"
"crypto/tls"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net"
"net/http"
"os"
"os/signal"
"path/filepath"
"strings"
"syscall"
"time"
"golang.org/x/term"
"vocat/internal/auth"
"vocat/internal/config"
"vocat/internal/developer"
"vocat/internal/device"
"vocat/internal/exportproxy"
"vocat/internal/extensions"
"vocat/internal/httpsmode"
"vocat/internal/loghub"
"vocat/internal/server"
"vocat/internal/store"
@@ -35,8 +44,25 @@ func main() {
args := os.Args[1:]
switch subcommand, rest := splitSubcommand(args); subcommand {
case "":
// No subcommand: run the server. Backward-compatible with the
// existing systemd unit (ExecStart=/opt/vocat/bin/vocat).
// No subcommand: TTY+root → interactive menu (operator on the host);
// otherwise run the server. systemd runs vocat with stdin=/dev/null
// (non-TTY) so the unit keeps starting the server unchanged. Non-root
// on a TTY also falls through to the server rather than erroring on
// runMenu's root requirement.
if term.IsTerminal(int(os.Stdin.Fd())) && os.Geteuid() == 0 {
if err := runMenu(logger); err != nil {
logger.Error("menu failed", "error", err)
os.Exit(1)
}
} else {
if err := run(logger, logs); err != nil {
logger.Error("server stopped", "error", err)
os.Exit(1)
}
}
case "serve":
// Explicit foreground server. Use this when vocat with no arguments
// would otherwise enter the menu (root on a TTY) but a server is wanted.
if err := run(logger, logs); err != nil {
logger.Error("server stopped", "error", err)
os.Exit(1)
@@ -53,6 +79,14 @@ func main() {
logger.Error("menu failed", "error", err)
os.Exit(1)
}
case "develop":
// Hidden subcommand: intentionally not listed in printUsage or the
// interactive menu. It toggles the developer-mode flag that gates the
// entire plugin/extension system; the flag takes effect on next start.
if err := runDevelop(rest, logger); err != nil {
logger.Error("develop failed", "error", err)
os.Exit(2)
}
case "help", "-h", "--help":
printUsage(os.Stdout)
default:
@@ -90,6 +124,50 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
return err
}
defer database.Close()
developerEnabled := isDeveloperEnabled(startupContext, database)
pluginRoot := filepath.Join(filepath.Dir(cfg.DatabasePath), "plugins")
legacyExportProxyConfig := filepath.Join(pluginRoot, exportproxy.ReservedID, "data", "configs.json")
if !developerEnabled {
if err := developer.ResetExperimental(startupContext, database); err != nil {
return fmt.Errorf("reset disabled developer settings: %w", err)
}
if err := exportproxy.RemoveLegacyConfig(legacyExportProxyConfig); err != nil {
return fmt.Errorf("remove legacy export proxy configuration: %w", err)
}
}
httpsManager, err := httpsmode.New(
startupContext,
database,
filepath.Join(filepath.Dir(cfg.DatabasePath), "tls"),
cfg.Address,
)
if err != nil {
return fmt.Errorf("configure self-signed HTTPS: %w", err)
}
// The plugin/extension system is gated behind a hidden developer-mode flag.
// When off (the default) the manager is never created and the server receives
// a nil Extensions handle, so every /extensions* and /plugin-assets/* route
// returns 503/404 and the SPA hides the plugin surface.
var extensionManager *extensions.Manager
var exportProxyManager *exportproxy.Manager
if developerEnabled {
exportProxyManager, err = exportproxy.New(startupContext, database, logger, legacyExportProxyConfig)
if err != nil {
return fmt.Errorf("create built-in export proxy: %w", err)
}
defer exportProxyManager.Close()
extensionManager, err = extensions.NewManager(
pluginRoot,
logger,
)
if err != nil {
return fmt.Errorf("create plugin manager: %w", err)
}
defer extensionManager.Close()
} else {
logger.Info("developer mode is off; plugin system disabled")
}
authService, err := auth.New(database, auth.Options{
SessionTTL: cfg.SessionTTL,
@@ -115,6 +193,7 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
if err := provisionDiscoveredDevices(startupContext, database, deviceManager); err != nil {
logger.Warn("automatic first-run device provisioning failed", "error", err)
}
restoreDefaultCellularRadios(startupContext, logger, database, deviceManager)
defer func() {
stopContext, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
@@ -125,7 +204,14 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
pollContext, cancelPolling := context.WithCancel(context.Background())
defer cancelPolling()
go pollDeviceSnapshots(pollContext, logger, database, deviceManager)
go restoreConfiguredCellularData(pollContext, logger, database, deviceManager)
go collectCellularTraffic(pollContext, logger, database)
go persistLogsToStore(pollContext, logger, logs, database)
if !developerEnabled {
go disableAllDeveloperCellularData(pollContext, logger, database, deviceManager)
} else {
go watchDeveloperDisable(pollContext, logger, database, deviceManager, exportProxyManager, legacyExportProxyConfig)
}
vowifiManager, err := configureVoWiFiRuntime(
startupContext,
@@ -154,6 +240,12 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
Logger: logger,
SecureCookies: cfg.SecureCookies,
MaxRequestBodyBytes: cfg.MaxRequestBodyBytes,
Extensions: extensionManager,
ExportProxy: exportProxyManager,
DeveloperEnabled: developerEnabled,
UpdateRepository: strings.TrimSpace(os.Getenv("VOCAT_REPO")),
UpdateToken: strings.TrimSpace(os.Getenv("GITHUB_TOKEN")),
HTTPS: httpsManager,
})
if err != nil {
return err
@@ -163,15 +255,35 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
handler.StartTelegramBot(pollContext)
handler.StartSMSNotificationDispatchers(pollContext)
httpServer := &http.Server{
Addr: cfg.Address,
Handler: handler,
ReadHeaderTimeout: 5 * time.Second,
ReadTimeout: 15 * time.Second,
WriteTimeout: 30 * time.Second,
IdleTimeout: 90 * time.Second,
MaxHeaderBytes: 1 << 20,
serverConfig := func(handler http.Handler) *http.Server {
return &http.Server{
Addr: cfg.Address,
Handler: handler,
ReadHeaderTimeout: 5 * time.Second,
ReadTimeout: 15 * time.Second,
WriteTimeout: 30 * time.Second,
IdleTimeout: 90 * time.Second,
MaxHeaderBytes: 1 << 20,
}
}
plainHandler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if httpsManager.Enabled() {
host := strings.TrimSpace(r.Host)
if host == "" {
host = cfg.Address
}
http.Redirect(w, r, "https://"+host+r.URL.RequestURI(), http.StatusPermanentRedirect)
return
}
handler.ServeHTTP(w, r)
})
plainServer := serverConfig(plainHandler)
tlsServer := serverConfig(handler)
baseListener, err := net.Listen("tcp", cfg.Address)
if err != nil {
return fmt.Errorf("listen on %s: %w", cfg.Address, err)
}
protocolMux := httpsmode.NewMultiplexer(baseListener, httpsManager)
signalContext, stopSignals := signal.NotifyContext(
context.Background(),
@@ -180,10 +292,17 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
)
defer stopSignals()
serverError := make(chan error, 1)
serverError := make(chan error, 2)
go func() {
logger.Info("HTTP server listening", "address", cfg.Address)
err := httpServer.ListenAndServe()
logger.Info("HTTP server listening", "address", cfg.Address, "self_signed_https", httpsManager.Enabled())
err := plainServer.Serve(protocolMux.Plain())
if errors.Is(err, http.ErrServerClosed) {
err = nil
}
serverError <- err
}()
go func() {
err := tlsServer.Serve(tls.NewListener(protocolMux.TLS(), httpsManager.TLSConfig()))
if errors.Is(err, http.ErrServerClosed) {
err = nil
}
@@ -192,21 +311,176 @@ func run(logger *slog.Logger, logs *loghub.Hub) error {
select {
case err := <-serverError:
_ = protocolMux.Close()
return err
case <-signalContext.Done():
logger.Info("shutdown signal received")
}
// Long-lived SSE and polling handlers use this context. Stop them before
// http.Server.Shutdown so they do not consume the entire graceful-shutdown
// deadline while waiting for a stream that is intentionally still active.
cancelPolling()
shutdownContext, cancelShutdown := context.WithTimeout(
context.Background(),
cfg.ShutdownTimeout,
)
defer cancelShutdown()
if err := httpServer.Shutdown(shutdownContext); err != nil {
_ = httpServer.Close()
return fmt.Errorf("graceful HTTP shutdown: %w", err)
shutdownErrors := make(chan error, 2)
go func() { shutdownErrors <- plainServer.Shutdown(shutdownContext) }()
go func() { shutdownErrors <- tlsServer.Shutdown(shutdownContext) }()
time.Sleep(10 * time.Millisecond)
_ = protocolMux.Close()
for range 2 {
if err := <-shutdownErrors; err != nil {
_ = plainServer.Close()
_ = tlsServer.Close()
return fmt.Errorf("graceful HTTP shutdown: %w", err)
}
}
return nil
}
// restoreDefaultCellularRadios repairs an interrupted VoWiFi teardown. CFUN=4
// survives process restarts, while the in-memory radio checkpoint does not. If
// VoWiFi is disabled and the current SIM has no explicit airplane policy, the
// automatic/default policy is cellular service and the modem must return to
// CFUN=1.
func restoreDefaultCellularRadios(
ctx context.Context,
logger *slog.Logger,
database *store.Store,
manager *device.Manager,
) {
configs, err := database.ListDevices(ctx)
if err != nil {
logger.Warn("startup cellular recovery: list devices", "error", err)
return
}
mapper := integration.ATMapper{Store: database, Devices: manager}
for _, config := range configs {
if config.VoWiFiEnabled {
continue
}
entry, err := mapper.Get(config.ID)
if err != nil || entry.Snapshot == nil || !entry.Snapshot.FlightMode {
continue
}
iccid := strings.TrimSpace(entry.Snapshot.ICCID)
if iccid != "" {
policy, policyErr := database.CardPolicy(ctx, iccid)
switch {
case policyErr == nil && policy.AirplaneEnabled:
continue
case policyErr != nil && !errors.Is(policyErr, store.ErrNotFound):
logger.Warn("startup cellular recovery: read card policy", "device_id", config.ID, "error", policyErr)
continue
}
}
restoreContext, cancel := context.WithTimeout(ctx, 10*time.Second)
_, err = manager.SetFlight(restoreContext, entry.ID, false)
cancel()
if err != nil {
logger.Warn("startup cellular recovery failed", "device_id", config.ID, "error", err)
continue
}
logger.Info("restored cellular radio after disabled VoWiFi", "device_id", config.ID, "iccid", iccid)
}
}
func restoreConfiguredCellularData(
ctx context.Context,
logger *slog.Logger,
database *store.Store,
manager *device.Manager,
) {
configs, err := database.ListDevices(ctx)
if err != nil {
logger.Warn("startup cellular data recovery: list devices", "error", err)
return
}
mapper := integration.ATMapper{Store: database, Devices: manager}
for _, config := range configs {
if !config.NetworkEnabled || config.VoWiFiEnabled {
continue
}
entry, err := mapper.Get(config.ID)
if err != nil {
continue
}
dataContext, cancel := context.WithTimeout(ctx, 60*time.Second)
_, err = manager.SetNetwork(dataContext, entry.ID, device.NetworkRequest{
Enabled: true, APN: config.APN, IPVersion: "IPV4V6",
})
cancel()
if err != nil {
logger.Warn("startup cellular data recovery failed", "device_id", config.ID, "error", err)
continue
}
logger.Info("restored protected cellular data route", "device_id", config.ID, "interface", config.Interface)
}
}
func disableAllDeveloperCellularData(
ctx context.Context,
logger *slog.Logger,
database *store.Store,
manager *device.Manager,
) {
configs, err := database.ListDevices(ctx)
if err != nil {
logger.Warn("developer cleanup: list devices", "error", err)
return
}
mapper := integration.ATMapper{Store: database, Devices: manager}
for _, config := range configs {
entry, err := mapper.Get(config.ID)
if err != nil {
continue
}
disableContext, cancel := context.WithTimeout(ctx, 30*time.Second)
_, err = manager.SetNetwork(disableContext, entry.ID, device.NetworkRequest{Enabled: false})
cancel()
if err != nil && ctx.Err() == nil {
logger.Warn("developer cleanup: stop cellular data", "device_id", config.ID, "error", err)
}
}
}
func watchDeveloperDisable(
ctx context.Context,
logger *slog.Logger,
database *store.Store,
manager *device.Manager,
exportProxy *exportproxy.Manager,
legacyConfigPath string,
) {
ticker := time.NewTicker(2 * time.Second)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
if developer.Enabled(ctx, database) {
continue
}
if exportProxy != nil {
if err := exportProxy.DeleteAllAndDisable(ctx); err != nil && ctx.Err() == nil {
logger.Warn("developer cleanup: delete export proxies", "error", err)
}
}
if err := exportproxy.RemoveLegacyConfig(legacyConfigPath); err != nil {
logger.Warn("developer cleanup: remove legacy export proxy configuration", "error", err)
}
if err := developer.ResetExperimental(ctx, database); err != nil && ctx.Err() == nil {
logger.Warn("developer cleanup: reset settings", "error", err)
}
disableAllDeveloperCellularData(ctx, logger, database, manager)
logger.Info("developer mode disabled; roaming data and export proxies were removed")
return
}
}
return <-serverError
}
func configureVoWiFiRuntime(
@@ -299,9 +573,20 @@ func newVoWiFiOrchestrator(
if message.Concat != nil && message.Concat.Total > 0 {
partsTotal = message.Concat.Total
}
messageID := message.MessageID
if message.Concat != nil && message.Concat.Total > 1 {
// A segment of a carrier-split long SMS over IMS. Address the whole
// message with a stable id so SaveSMSMessage folds every segment
// into one progressively merged row instead of one row per segment.
messageID = store.StableConcatMessageID(
"ims", deviceConfig.ModemIMEI, message.DeviceID, message.From,
message.Concat.Reference, message.Concat.Total,
)
}
_, saveErr := database.SaveSMSMessage(ctx, store.SMSMessage{
MessageID: message.MessageID,
MessageID: messageID,
DeviceID: message.DeviceID,
ModemIMEI: deviceConfig.ModemIMEI,
IMSI: message.IMSI,
Peer: message.From,
Direction: "inbound",
@@ -318,6 +603,7 @@ func newVoWiFiOrchestrator(
OnSMSStatus: func(ctx context.Context, report ims.ReceivedSMSStatus) error {
deliveryReport := store.SMSDeliveryReport{
DeviceID: report.DeviceID,
ModemIMEI: deviceConfig.ModemIMEI,
IMSI: report.IMSI,
Peer: report.To,
Source: "ims",
+192 -43
View File
@@ -3,6 +3,7 @@ package main
import (
"bufio"
"context"
"encoding/json"
"errors"
"fmt"
"log/slog"
@@ -16,18 +17,62 @@ import (
"vocat/internal/auth"
"vocat/internal/config"
"vocat/internal/store"
"vocat/internal/update"
)
//envFilePath is the systemd EnvironmentFile that carries VOCAT_ADMIN_PASSWORD.
// envFilePath is the systemd EnvironmentFile that carries VOCAT_ADMIN_PASSWORD.
// EnsureAdmin reseeds the DB from it on every start, so change-password must
// rewrite it or the next restart reverts the password.
const envFilePath = "/etc/vocat/env"
const systemdUnitPath = "/etc/systemd/system/vocat.service"
// runMenu is the interactive lifecycle menu: change password, restart the
// systemd unit, or fully uninstall vocat. It must run as root on the host
// (needs systemctl + the 0600 env file). Docker deployments do not use it.
// defaultDatabasePath is the install-default SQLite location written into the
// systemd unit by scripts/install.sh. Used only when VOCAT_DATABASE_PATH is not
// already set in the operator's environment.
const defaultDatabasePath = "/opt/vocat/data/vocat.db"
// uiPreferencesSettingKey is the same app_settings key the Web UI's
// /api/settings/preferences handler reads and writes (see general_api.go). The
// menu toggles language through it so a single preference is shared with the
// SPA and the backend i18n layer.
const uiPreferencesSettingKey = "ui.preferences"
// loadMenuEnv ensures the menu reaches the production config that systemd
// would otherwise inject. When an operator runs `sudo vocat` on the host, the
// shell has not sourced /etc/vocat/env (a systemd EnvironmentFile, not a shell
// rc) and VOCAT_DATABASE_PATH is unset, so config.Load() would resolve a
// CWD-relative ./data/vocat.db — a different, empty database than
// /opt/vocat/data/vocat.db the service uses. This loads the installed env file
// for VOCAT_ADMIN_PASSWORD and pins VOCAT_DATABASE_PATH to the install default,
// without overriding any value the operator already exported.
func loadMenuEnv() {
if _, ok := os.LookupEnv("VOCAT_DATABASE_PATH"); !ok {
_ = os.Setenv("VOCAT_DATABASE_PATH", defaultDatabasePath)
}
if data, err := os.ReadFile(envFilePath); err == nil {
for _, line := range strings.Split(string(data), "\n") {
line = strings.TrimSpace(line)
if line == "" || strings.HasPrefix(line, "#") {
continue
}
eq := strings.IndexByte(line, '=')
if eq < 0 {
continue
}
key := strings.TrimSpace(line[:eq])
val := strings.TrimSpace(line[eq+1:])
if _, ok := os.LookupEnv(key); !ok {
_ = os.Setenv(key, val)
}
}
}
}
// runMenu is the interactive lifecycle menu: toggle language, change password,
// restart the systemd unit, self-update, or fully uninstall vocat. It must run
// as root on the host (needs systemctl + the 0600 env file). Docker deployments
// do not use it.
func runMenu(logger *slog.Logger) error {
if os.Geteuid() != 0 {
return errors.New("vocat menu must run as root (needs systemctl and /etc/vocat/env)")
@@ -37,8 +82,13 @@ func runMenu(logger *slog.Logger) error {
return errors.New("vocat menu requires an interactive terminal")
}
lang := promptLanguage()
loadMenuEnv()
lang, langErr := loadMenuLanguage()
menu := newMenu(lang)
if langErr != nil {
logger.Warn("menu: load language preference failed; defaulting to English", "error", langErr)
}
reader := bufio.NewReader(os.Stdin)
for {
@@ -55,44 +105,69 @@ func runMenu(logger *slog.Logger) error {
choice := strings.TrimSpace(line)
switch choice {
case "1":
if err := menuChangePassword(reader, menu, logger); err != nil {
if err := menuToggleLanguage(menu, logger); err != nil {
fmt.Println(menu.errorPrefix(err))
}
case "2":
if err := menuRestart(menu); err != nil {
if err := menuChangePassword(reader, menu, logger); err != nil {
fmt.Println(menu.errorPrefix(err))
}
case "3":
if err := menuUninstall(reader, menu); err != nil {
if err := menuRestart(menu); err != nil {
fmt.Println(menu.errorPrefix(err))
}
case "0", "":
fmt.Println(menu.bye())
return nil
case "4":
if err := menuUpdate(menu, logger); err != nil {
fmt.Println(menu.errorPrefix(err))
}
case "0":
if err := menuUninstall(reader, menu); err != nil {
fmt.Println(menu.errorPrefix(err))
} else {
return nil
}
default:
fmt.Println(menu.invalid())
}
}
}
// promptLanguage asks for 中文 (1) or English (2) once per invocation. The
// user chose to re-ask every run rather than persist a language preference.
func promptLanguage() string {
reader := bufio.NewReader(os.Stdin)
for {
fmt.Println("选择语言 / Select language: 1) 中文 2) English")
fmt.Print("> ")
line, err := reader.ReadString('\n')
if err != nil {
return "zh"
}
switch strings.TrimSpace(line) {
case "1", "":
return "zh"
case "2":
return "en"
}
// loadMenuLanguage reads the persisted UI preference (the same ui.preferences
// app_setting the Web UI writes) and returns "zh" or "en". When no record
// exists yet it returns "en", matching the Web default in writeUIPreferences.
// Any failure is surfaced to the caller, which logs a warning and keeps the
// default rather than blocking the menu.
func loadMenuLanguage() (string, error) {
cfg, err := config.Load()
if err != nil {
return "en", fmt.Errorf("%w: %v", errMenuConfig, err)
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
database, err := store.Open(ctx, cfg.DatabasePath)
if err != nil {
return "en", fmt.Errorf("%w: %v", errMenuStore, err)
}
defer database.Close()
setting, err := database.AppSetting(ctx, uiPreferencesSettingKey)
if err != nil {
if errors.Is(err, store.ErrNotFound) {
return "en", nil
}
return "en", fmt.Errorf("%w: %v", errMenuStore, err)
}
var prefs struct {
Language string `json:"language"`
}
if err := json.Unmarshal(setting.Value, &prefs); err != nil {
return "en", nil
}
if prefs.Language == "zh" {
return "zh", nil
}
return "en", nil
}
func menuChangePassword(reader *bufio.Reader, m *menu, logger *slog.Logger) error {
@@ -212,6 +287,45 @@ func rewriteEnvPassword(newPassword string) error {
return os.Rename(tmpName, envFilePath)
}
// menuToggleLanguage flips the persisted language preference between "zh" and
// "en" by writing the same ui.preferences app_setting the Web UI uses, then
// switches the menu's own language so the next prompt renders in the new
// language. The Web SPA picks up the change on its next preferences fetch; the
// menu never needs to call i18n.Set itself since it carries its own lang copy.
func menuToggleLanguage(m *menu, logger *slog.Logger) error {
cfg, err := config.Load()
if err != nil {
return fmt.Errorf("%w: %v", errMenuConfig, err)
}
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
defer cancel()
database, err := store.Open(ctx, cfg.DatabasePath)
if err != nil {
return fmt.Errorf("%w: %v", errMenuStore, err)
}
defer database.Close()
current := m.lang
next := "zh"
if current == "zh" {
next = "en"
}
payload, err := json.Marshal(map[string]string{"language": next})
if err != nil {
return fmt.Errorf("%w: %v", errMenuStore, err)
}
if err := database.UpsertAppSetting(ctx, store.AppSetting{
Key: uiPreferencesSettingKey,
Value: payload,
}); err != nil {
return fmt.Errorf("%w: %v", errMenuStore, err)
}
m.lang = next
fmt.Println(m.languageSwitched())
return nil
}
func menuRestart(m *menu) error {
if _, err := exec.LookPath("systemctl"); err != nil {
return errNoSystemctl
@@ -224,6 +338,23 @@ func menuRestart(m *menu) error {
return nil
}
// menuUpdate delegates to the self-updater in internal/update. It resolves the
// repo the same way update.Run itself does ($VOCAT_REPO or the default) and lets
// that package handle the check/download/verify/replace/restart flow. The
// running menu process keeps the old binary until the operator exits; only the
// systemd service runs the new build after restartService.
func menuUpdate(m *menu, logger *slog.Logger) error {
repo := strings.TrimSpace(os.Getenv("VOCAT_REPO"))
if repo == "" {
repo = update.DefaultRepository
}
fmt.Println(m.updateChecking())
if err := update.Run(logger, []string{"--repo", repo}); err != nil {
return fmt.Errorf("%w: %v", errUpdateFailed, err)
}
return nil
}
// menuUninstall performs full removal: stop/disable the unit, delete the unit,
// remove /opt/vocat (binary + data + SQLite DB), remove the env file, reload
// systemd, and best-effort delete the vocat user.
@@ -261,6 +392,7 @@ var (
errPasswordsDiffer = errors.New("menu: passwords do not match")
errNoSystemctl = errors.New("menu: systemctl not found")
errRestartFailed = errors.New("menu: restart failed")
errUpdateFailed = errors.New("menu: update failed")
errMenuConfig = errors.New("menu: load configuration")
errMenuStore = errors.New("menu: open database")
errMenuAuth = errors.New("menu: auth service")
@@ -277,19 +409,24 @@ func newMenu(lang string) *menu { return &menu{lang: lang} }
func (m *menu) msg(key string) string {
const zh, en = 0, 1
table := map[string][2]string{
"title": {"vocat 管理菜单", "vocat management menu"},
"opt_change": {"1) 修改密码", "1) Change password"},
"opt_restart": {"2) 重启服务", "2) Restart service"},
"opt_uninstall": {"3) 卸载程序", "3) Uninstall"},
"opt_exit": {"0) 退出", "0) Exit"},
"prompt": {"请选择: ", "Select: "},
"invalid": {"无效选项,请重试。", "Invalid choice, try again."},
"bye": {"再见。", "Bye."},
"cur_pw": {"当前密码: ", "Current password: "},
"new_pw": {"新密码 (至少 12 位): ", "New password (min 12 chars): "},
"confirm_pw": {"确认新密码: ", "Confirm new password: "},
"pw_changed": {"密码已修改。重启后仍然有效。", "Password changed. Survives restart."},
"restarted": {"服务已重启。", "Service restarted."},
"title": {"vocat 管理菜单", "vocat management menu"},
"opt_lang": {"1) 切换中英文", "1) Toggle language"},
"opt_change": {"2) 修改账号密码", "2) Change admin password"},
"opt_restart": {"3) 重启软件", "3) Restart software"},
"opt_update": {"4) 更新软件", "4) Update software"},
"opt_uninstall": {"0) 卸载软件", "0) Uninstall software"},
"prompt": {"请选择: ", "Select: "},
"invalid": {"无效选项,请重试。按 Ctrl+C 退出。", "Invalid choice, try again. Press Ctrl+C to exit."},
"cur_pw": {"当前密码: ", "Current password: "},
"new_pw": {"新密码 (至少 12 位): ", "New password (min 12 chars): "},
"confirm_pw": {"确认新密码: ", "Confirm new password: "},
"pw_changed": {"密码已修改。重启后仍然有效。", "Password changed. Survives restart."},
"lang_switched": {
"语言已切换。Web 界面下次刷新后同步。",
"Language switched. The web UI syncs on next refresh.",
},
"upd_checking": {"正在检查更新…", "Checking for updates…"},
"restarted": {"软件已重启。", "Software restarted."},
"uninstall_warn": {
"警告: 将删除程序、数据与配置,且不可恢复!",
"WARNING: removes the program, data and config. Irreversible!",
@@ -311,11 +448,12 @@ func (m *menu) msg(key string) string {
func (m *menu) title() string { return m.msg("title") }
func (m *menu) prompt() string { return m.msg("prompt") }
func (m *menu) invalid() string { return m.msg("invalid") }
func (m *menu) bye() string { return m.msg("bye") }
func (m *menu) currentPassword() string { return m.msg("cur_pw") }
func (m *menu) newPassword() string { return m.msg("new_pw") }
func (m *menu) confirmPassword() string { return m.msg("confirm_pw") }
func (m *menu) passwordChanged() string { return m.msg("pw_changed") }
func (m *menu) languageSwitched() string { return m.msg("lang_switched") }
func (m *menu) updateChecking() string { return m.msg("upd_checking") }
func (m *menu) restarted() string { return m.msg("restarted") }
func (m *menu) uninstallWarn() string { return m.msg("uninstall_warn") }
func (m *menu) uninstallConfirm() string { return m.msg("uninstall_confirm") }
@@ -323,7 +461,13 @@ func (m *menu) uninstallCancelled() string { return m.msg("uninstall_cancelled")
func (m *menu) uninstalled() string { return m.msg("uninstalled") }
func (m *menu) options() []string {
return []string{m.msg("opt_change"), m.msg("opt_restart"), m.msg("opt_uninstall"), m.msg("opt_exit")}
return []string{
m.msg("opt_lang"),
m.msg("opt_change"),
m.msg("opt_restart"),
m.msg("opt_update"),
m.msg("opt_uninstall"),
}
}
func (m *menu) errorPrefix(err error) string {
@@ -348,6 +492,11 @@ func (m *menu) errorPrefix(err error) string {
return "Restart failed."
}
return "重启失败。"
case errors.Is(err, errUpdateFailed):
if m.lang == "en" {
return "Update failed."
}
return "更新失败。"
case errors.Is(err, errMenuConfig):
if m.lang == "en" {
return "Failed to load configuration."
+123
View File
@@ -0,0 +1,123 @@
package main
import (
"context"
"log/slog"
"math"
"strings"
"time"
"vocat/internal/store"
)
const cellularTrafficSampleInterval = 30 * time.Second
type interfaceTrafficSample struct {
interfaceName string
rxBytes uint64
txBytes uint64
}
func collectCellularTraffic(ctx context.Context, logger *slog.Logger, database *store.Store) {
previous := make(map[string]interfaceTrafficSample)
var lastPrune time.Time
collect := func() {
now := time.Now()
if lastPrune.IsZero() || now.Sub(lastPrune) >= 24*time.Hour {
lastPrune = now
if _, err := database.DeleteTrafficBefore(ctx, now.Add(-35*24*time.Hour)); err != nil && ctx.Err() == nil {
logger.Warn("prune old cellular traffic", "error", err)
}
}
configs, err := database.ListDevices(ctx)
if err != nil {
if ctx.Err() == nil {
logger.Warn("list devices for cellular traffic collection", "error", err)
}
return
}
active := make(map[string]struct{}, len(configs))
for _, config := range configs {
interfaceName := strings.TrimSpace(config.Interface)
if !config.NetworkEnabled || interfaceName == "" {
delete(previous, config.ID)
continue
}
active[config.ID] = struct{}{}
rxBytes, txBytes, err := readInterfaceTrafficCounters(interfaceName)
if err != nil {
// Interfaces can briefly disappear while QMI reconnects. The next
// successful read establishes a fresh baseline, so no reconnect
// traffic is accidentally counted twice.
delete(previous, config.ID)
continue
}
rxDelta, txDelta, ok := trafficCounterDelta(previous[config.ID], interfaceName, rxBytes, txBytes)
previous[config.ID] = interfaceTrafficSample{
interfaceName: interfaceName,
rxBytes: rxBytes,
txBytes: txBytes,
}
if !ok || (rxDelta == 0 && txDelta == 0) {
continue
}
for bucket, periodStart := range trafficBucketPeriods(time.Now()) {
if err := database.AddTrafficBucket(ctx, store.TrafficBucket{
DeviceID: config.ID,
Bucket: bucket,
PeriodStart: periodStart,
RXBytes: rxDelta,
TXBytes: txDelta,
}); err != nil && ctx.Err() == nil {
logger.Warn("record cellular traffic", "device", config.ID, "bucket", bucket, "error", err)
}
}
}
for deviceID := range previous {
if _, ok := active[deviceID]; !ok {
delete(previous, deviceID)
}
}
}
collect()
ticker := time.NewTicker(cellularTrafficSampleInterval)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return
case <-ticker.C:
collect()
}
}
}
func trafficCounterDelta(previous interfaceTrafficSample, interfaceName string, rxBytes, txBytes uint64) (int64, int64, bool) {
if previous.interfaceName == "" || previous.interfaceName != interfaceName || rxBytes < previous.rxBytes || txBytes < previous.txBytes {
return 0, 0, false
}
rxDelta := rxBytes - previous.rxBytes
txDelta := txBytes - previous.txBytes
if rxDelta > math.MaxInt64 || txDelta > math.MaxInt64 {
return 0, 0, false
}
return int64(rxDelta), int64(txDelta), true
}
func trafficBucketPeriods(now time.Time) map[string]time.Time {
local := now.In(time.Local)
year, month, day := local.Date()
dayStart := time.Date(year, month, day, 0, 0, 0, 0, time.Local).UTC()
return map[string]time.Time{
"hour": now.UTC().Truncate(time.Minute),
"day": now.UTC().Truncate(time.Hour),
"week": dayStart,
"month": dayStart,
}
}
+39
View File
@@ -0,0 +1,39 @@
//go:build linux
package main
import (
"fmt"
"net"
"os"
"path/filepath"
"strconv"
"strings"
)
func readInterfaceTrafficCounters(interfaceName string) (uint64, uint64, error) {
iface, err := net.InterfaceByName(interfaceName)
if err != nil {
return 0, 0, err
}
read := func(counter string) (uint64, error) {
value, err := os.ReadFile(filepath.Join("/sys/class/net", iface.Name, "statistics", counter))
if err != nil {
return 0, err
}
parsed, err := strconv.ParseUint(strings.TrimSpace(string(value)), 10, 64)
if err != nil {
return 0, fmt.Errorf("parse %s %s counter: %w", iface.Name, counter, err)
}
return parsed, nil
}
rxBytes, err := read("rx_bytes")
if err != nil {
return 0, 0, err
}
txBytes, err := read("tx_bytes")
if err != nil {
return 0, 0, err
}
return rxBytes, txBytes, nil
}
+9
View File
@@ -0,0 +1,9 @@
//go:build !linux
package main
import "errors"
func readInterfaceTrafficCounters(string) (uint64, uint64, error) {
return 0, 0, errors.New("interface traffic counters are only available on Linux")
}
+38
View File
@@ -0,0 +1,38 @@
package main
import (
"testing"
"time"
)
func TestTrafficCounterDelta(t *testing.T) {
previous := interfaceTrafficSample{interfaceName: "wwan0", rxBytes: 100, txBytes: 50}
rx, tx, ok := trafficCounterDelta(previous, "wwan0", 175, 90)
if !ok || rx != 75 || tx != 40 {
t.Fatalf("delta = (%d, %d, %v), want (75, 40, true)", rx, tx, ok)
}
if _, _, ok := trafficCounterDelta(previous, "wwan1", 175, 90); ok {
t.Fatal("interface change must establish a new baseline")
}
if _, _, ok := trafficCounterDelta(previous, "wwan0", 90, 40); ok {
t.Fatal("counter reset must establish a new baseline")
}
}
func TestTrafficBucketPeriods(t *testing.T) {
now := time.Date(2026, 8, 10, 12, 34, 56, 0, time.Local)
periods := trafficBucketPeriods(now)
if got := periods["hour"]; !got.Equal(now.UTC().Truncate(time.Minute)) {
t.Fatalf("hour period = %s", got)
}
if got := periods["day"]; !got.Equal(now.UTC().Truncate(time.Hour)) {
t.Fatalf("day period = %s", got)
}
localDay := periods["week"].In(time.Local)
if localDay.Hour() != 0 || localDay.Minute() != 0 || localDay.Day() != 10 {
t.Fatalf("week period = %s, want local day start", periods["week"])
}
if !periods["month"].Equal(periods["week"]) {
t.Fatal("week and month should share daily periods")
}
}
+1 -1
View File
@@ -29,7 +29,7 @@ ProtectKernelModules=true
ProtectKernelTunables=true
ProtectControlGroups=true
ReadWritePaths=/opt/vocat/data
RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6 AF_NETLINK
RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6 AF_NETLINK AF_PACKET
RestrictRealtime=true
LockPersonality=true
MemoryDenyWriteExecute=true
+60
View File
@@ -0,0 +1,60 @@
# VoCat Docker Compose deployment.
#
# First-time setup:
# cp .env.example .env # then edit VOCAT_ADMIN_PASSWORD
# docker compose pull # fetch the prebuilt GHCR image
# docker compose up -d # start
#
# Build locally from this repo instead of using the GHCR image:
# docker compose up -d --build
#
# In-container binary self-update is intentionally disabled (VOCAT_CONTAINER=docker
# makes the server return 409 on the apply endpoint). Update by pulling a new
# image and recreating the container:
# docker compose pull && docker compose up -d
services:
vocat:
# Use the prebuilt multi-arch image from GHCR. Override with
# --build to compile from the local Dockerfile instead.
image: ghcr.io/mengmengcode/vocat:latest
pull_policy: missing
build:
context: .
dockerfile: Dockerfile
container_name: vocat
restart: unless-stopped
# Host network mode: the export-proxy plugin uses SO_BINDTODEVICE to pin
# outbound proxy traffic to the modem interface (wwan0) so roaming data
# egresses only the module — never the host's default route. That syscall
# needs the host network namespace visible inside the container, which
# network_mode: host provides directly. Port publishing is therefore
# meaningless (the container shares the host stack and vocat binds
# 0.0.0.0:7575 itself); proxy ports opened by the plugin are likewise
# reachable on the host IP without explicit mapping.
network_mode: host
# VoWiFi / eSIM / IMS paths need raw sockets (IPsec, netlink). The systemd
# unit grants CAP_NET_ADMIN + CAP_NET_RAW; mirror that here.
cap_add:
- NET_ADMIN
- NET_RAW
environment:
# Marks the process as containerized: the web UI then advertises
# "pull new image" instead of attempting an in-place binary update.
VOCAT_CONTAINER: docker
# VOCAT_ADDR / VOCAT_DATABASE_PATH are set in the Dockerfile; override
# only if you want non-default values. Sensitive values come from .env.
VOCAT_ADMIN_PASSWORD: ${VOCAT_ADMIN_PASSWORD:?set VOCAT_ADMIN_PASSWORD in .env}
volumes:
# SQLite database + persistent state. Named volume (not a bind mount)
# because the container runs as uid 1000 (vocat) while a bind-mounted
# host dir would be root-owned and unwritable. Docker gives the named
# volume the image's uid 1000 ownership automatically.
- vocat-data:/opt/vocat/data
volumes:
vocat-data:
+295
View File
@@ -0,0 +1,295 @@
<p align="center">
<img src="../web/public/favicon.svg" width="96" alt="Vocat">
</p>
<h1 align="center">VoCat</h1>
<p align="center">
<img alt="Go" src="https://img.shields.io/badge/Go-1.25-00ADD8?style=flat-square&logo=go&logoColor=white">
<img alt="React" src="https://img.shields.io/badge/React-19-61DAFB?style=flat-square&logo=react&logoColor=111111">
<img alt="TypeScript" src="https://img.shields.io/badge/TypeScript-5.8-3178C6?style=flat-square&logo=typescript&logoColor=white">
<img alt="Vite" src="https://img.shields.io/badge/Vite-7-646CFF?style=flat-square&logo=vite&logoColor=white">
<img alt="Tailwind CSS" src="https://img.shields.io/badge/Tailwind_CSS-3-06B6D4?style=flat-square&logo=tailwindcss&logoColor=white">
<img alt="SQLite" src="https://img.shields.io/badge/SQLite-Embedded-003B57?style=flat-square&logo=sqlite&logoColor=white">
</p>
<p align="center">
<img alt="Linux" src="https://img.shields.io/badge/Linux-amd64_%7C_386_%7C_arm64_%7C_armv7-FCC624?style=flat-square&logo=linux&logoColor=111111">
<img alt="Docker" src="https://img.shields.io/badge/Docker-Multi--Arch-2496ED?style=flat-square&logo=docker&logoColor=white">
<img alt="WiFi Calling" src="https://img.shields.io/badge/WiFi_Calling-IMS_SMS-7B1FA2?style=flat-square">
<img alt="eSIM" src="https://img.shields.io/badge/eSIM-LPA_%2F_eUICC-009688?style=flat-square">
<img alt="Telegram" src="https://img.shields.io/badge/Telegram-Bot-26A5E4?style=flat-square&logo=telegram&logoColor=white">
<img alt="GitHub Actions" src="https://img.shields.io/badge/GitHub_Actions-Release-2088FF?style=flat-square&logo=githubactions&logoColor=white">
</p>
[English](../README.md) | **简体中文**
Vocat 是一款面向 Quectel EC20/EC25 系列蜂窝模组的开源 Web 控制面板与工程工具套件。它在一个自包含的服务中整合了模组发现、实时射频状态、AT 与 USSD 终端、短信、WiFi Calling(WiFi 通话)、eSIM 管理、网络选择、代理路由、通知、审计日志以及发布自动化。
后端使用 Go 编写,界面采用 React 与 TypeScript 构建,生产环境前端被嵌入进 Go 二进制中。单个可执行文件即包含完整的 Web 应用,并使用 SQLite 进行持久化存储。
<p align="center">
<img src="../img/image.png">
<img src="../img/image-1.png">
</p>
## 功能
| 领域 | Vocat 提供的能力 |
| --- | --- |
| 设备管理 | 自动串口/USB 发现、多模组支持、设备友好名称、概览实时刷新、模组重启、飞行模式以及 USB 网卡模式控制。 |
| 射频与网络 | 注册状态、运营商、信号指标、RSRP/RSRQ/SINR、网络模式、频段、信道、运营商扫描以及自动/手动选网。 |
| AT 与 USSD | 交互式 AT 终端、命令历史、原始模组响应、USSD 发起/继续/取消流程以及清晰的模组错误上报。 |
| 短信 | 蜂窝与 IMS 短信直接发送、入站同步、长短信合并、送达报告、会话历史、未读状态、时间戳以及逐条消息的送达状态。 |
| WiFi Calling | IKEv2/ePDG 隧道建立、EAP-AKA 鉴权、IMS 注册、IMS 短信、重连控制、状态诊断以及按设备路由。 |
| eSIM 与 eUICC | eUICC 发现、EID 与生产信息、证书元数据、多 eUICC 清单、已安装配置文件列表、启用/禁用/切换操作,以及在卡片支持时进行下载、重命名和删除。 |
| 卡策略 | 基于 ICCID 的 WiFi Calling 与飞行模式行为,策略即时应用。 |
| 代理路由 | 上游 SOCKS 路由、设备绑定、国家规则、TCP 可达性检查以及面向 WiFi Calling 数据路径的 UDP Associate 检查。 |
| 通知 | 通过 Telegram、Bark、邮件、Pushplus 以及签名 Webhook 转发新入站短信,每条短信单独推送。 |
| Telegram 机器人 | 设备状态、已安装配置文件列表与切换、WiFi Calling 控制、短信发送、定时拨号并自动挂断、通话状态、接听与挂断命令。敏感操作需要管理员确认。 |
| 运维 | 鉴权、CSRF 防护、访问策略、审计事件、实时日志、日志留存、健康检查、响应式布局、深色模式以及中英文应用界面。 |
| 分发 | 静态 Linux 二进制、systemd 安装脚本、带 SHA-256 校验的自更新、Docker 镜像、GHCR 发布以及 GitHub Actions 发布构建。 |
## 支持的硬件
Vocat 面向基于高通芯片、并暴露兼容 AT、QMI、串口与 USB 网络接口的 Quectel 模组,包括:
- Quectel EC20
- Quectel EC25
- Quectel EG25 系列
- 兼容的 EG600 及相关模组
可用功能取决于模组固件、USB 复合设备配置、SIM/eSIM 能力、主机驱动、无线网络以及运营商配置。
## 安装
### Linux 一键安装
```bash
curl -fsSL https://raw.githubusercontent.com/MengMengCode/VoCat/master/scripts/install.sh | sudo bash
```
安装指定版本:
```bash
curl -fsSL https://raw.githubusercontent.com/MengMengCode/VoCat/master/scripts/install.sh -o install.sh
sudo bash install.sh 0.0.2
```
安装程序会:
- 检测 `amd64``386``arm64``armv7` 架构;
- 下载对应的 GitHub Release 二进制;
- 对照 `SHA256SUMS` 进行校验;
- 将 Vocat 安装到 `/opt/vocat`;
- 创建具有 Vocat 所需硬件与网络访问权限的强化版 systemd 服务;
- 将运行时配置存放在 `/etc/vocat/env`;
- 首次安装时生成随机初始管理员密码。
安装完成后打开:
```text
http://<服务器地址>:7575
```
### 手动二进制安装
从 GitHub Releases 下载对应的二进制与 `SHA256SUMS`:
| 平台 | 发布文件 |
| --- | --- |
| Linux x86-64 | `vocat-linux-amd64` |
| Linux x86 32 位 | `vocat-linux-386` |
| Linux ARM64 | `vocat-linux-arm64` |
| Linux ARMv7 | `vocat-linux-armv7` |
校验并安装:
```bash
sha256sum -c SHA256SUMS --ignore-missing
sudo install -d -m 0755 /opt/vocat/bin /opt/vocat/data
sudo install -m 0755 vocat-linux-amd64 /opt/vocat/bin/vocat
sudo env \
VOCAT_DATABASE_PATH=/opt/vocat/data/vocat.db \
VOCAT_ADMIN_PASSWORD=change-this-password \
/opt/vocat/bin/vocat serve
```
该手动命令会在前台运行 Vocat。请使用 `vocat serve` 以直接启动服务器;在 TTY 下以 root 运行无参数的 `vocat` 会进入交互式管理菜单。如需托管的 systemd 服务与自动重启,请使用一键安装脚本。
### Docker
如果 Linux 主机需要发现每一个接入的受支持 Quectel 模组,并持续感知 USB 热插拔事件,请以硬件访问模式运行 Vocat:
```bash
docker pull ghcr.io/mengmengcode/vocat:latest
docker run -d \
--name vocat \
--restart unless-stopped \
--network host \
--privileged \
--user 0:0 \
-e VOCAT_ADMIN_PASSWORD=change-this-password \
-v vocat-data:/opt/vocat/data \
-v /dev:/dev \
-v /sys:/sys:ro \
ghcr.io/mengmengcode/vocat:latest
```
容器启动后打开 `http://<服务器地址>:7575`。主机网络是必需的,这样 QMI 网络接口才能对 Vocat 可见;而特权设备访问是串口、QMI 控制节点、TUN 接口、网络配置以及容器启动后新增设备所必需的。`/dev` 挂载使新的 `ttyUSB*``ttyACM*``cdc-wdm*` 节点无需重建容器即可见。
该模式有意赋予 Vocat 对主机设备与网络栈的广泛访问权限,仅在受信任的 Linux 主机上使用。自动发现目前仅识别受支持的 Quectel USB 模组(USB 厂商 ID `2c7c`),不识别任意品牌的模组。仅用 `--device` 映射单个节点(例如 `/dev/ttyUSB2``/dev/cdc-wdm0`)会将容器限定在这些固定节点上,无法提供完整的多设备或热插拔发现。
GHCR 镜像发布为 `linux/amd64``linux/arm64`
## 配置
Vocat 先从 `VOCAT_CONFIG` 读取可选的 JSON 配置文件,再应用 `VOCAT_*` 环境变量。环境变量优先级更高。
| 环境变量 | 默认值 | 说明 |
| --- | --- | --- |
| `VOCAT_ADDR` | `0.0.0.0:7575` | HTTP 监听地址。 |
| `VOCAT_DATABASE_PATH` | `./data/vocat.db` | SQLite 数据库路径。 |
| `VOCAT_ADMIN_USERNAME` | `admin` | 初始管理员用户名。 |
| `VOCAT_ADMIN_PASSWORD` | `admin` | 初始管理员密码。暴露服务前请务必修改。 |
| `VOCAT_SESSION_TTL` | `24h` | 鉴权会话有效期。 |
| `VOCAT_SECURE_COOKIES` | `false` | 在使用 HTTPS 时将会话 Cookie 标记为安全。 |
| `VOCAT_SHUTDOWN_TIMEOUT` | `10s` | 优雅关闭超时时间。 |
| `VOCAT_MAX_REQUEST_BODY_BYTES` | `1048576` | API 请求体最大字节数。 |
| `VOCAT_REPO` | `MengMengCode/VoCat` | 自更新器使用的受信任 GitHub 仓库,格式为 `owner/name`。 |
| `GITHUB_TOKEN` | 空 | 可选的 GitHub token,用于私有仓库或更高的 API 限额。 |
请勿将 Telegram token、SMTP 密码、Webhook 密钥、SIM 凭据或其他私密数据存放在仓库中。请通过应用设置或受保护的环境文件来配置它们。
## Telegram 机器人
启用 Telegram 通知并配置好 Chat ID 与 Admin ID 后,机器人支持:
```text
/status [设备]
/esim <设备>
/switch <设备> <iccid>
/wfc <设备> <status|on|off|reconnect>
/sms <设备> <号码> <内容>
/call <设备> <号码> <秒数>
/calls <设备>
/answer <设备>
/hangup <设备>
```
配置文件切换、短信提交与拨号使用一次性确认按钮。定时拨号会执行模组拨号动作,并在 1–600 秒后自动挂断;不会捕获或处理通话音频。机器人不暴露 eSIM 下载、删除或重命名命令。
## 更新
检查是否有更新的 GitHub Release:
```bash
vocat update --check --repo MengMengCode/VoCat
```
安装最新发布版:
```bash
sudo vocat update --repo MengMengCode/VoCat
```
更新器会下载与当前 Linux 架构匹配的二进制,使用已发布的 `SHA256SUMS` 进行校验,原子性地替换可执行文件,并在可用时重启 `vocat` systemd 服务。
Docker 安装的更新方式:
```bash
docker pull ghcr.io/mengmengcode/vocat:latest
```
拉取新镜像后重建容器。
## 开发
依赖要求:
- Go 1.25 或更新版本
- Node.js 20 或更新版本
- npm
运行前端开发服务器:
```bash
cd web
npm install
npm run dev
```
构建嵌入的前端并启动后端:
```bash
cd web
npm run build
cd ..
go run ./cmd/vocat
```
运行全部测试:
```bash
go test ./...
```
构建生产二进制:
```bash
go build -trimpath -ldflags "-s -w" -o vocat ./cmd/vocat
```
## 发布自动化
推送版本标签会触发两个 GitHub Actions 工作流:
- `release-binaries` 构建并发布 `amd64``386``arm64``armv7` 二进制及 `SHA256SUMS`
- `docker` 构建并向 GitHub Container Registry 发布多架构镜像。
```bash
git tag v0.2.0
git push origin v0.2.0
```
## 项目结构
```text
cmd/vocat/ 应用入口与 CLI
internal/device/ 模组发现与设备控制
internal/modem/ AT 会话与响应处理
internal/server/ HTTP API、通知与内嵌 Web 服务器
internal/store/ SQLite 持久化
internal/update/ GitHub Release 自更新器
internal/vowifi/ IKE、EAP-AKA、IMS 与 WiFi Calling 运行时
scripts/install.sh Linux 安装与更新脚本
web/src/ React 与 TypeScript 前端
.github/workflows/ 二进制与 Docker 发布自动化
```
## 合规使用
蜂窝模组与 eSIM 操作可能影响用户服务、已存储的配置文件、网络注册以及硬件状态。请做好备份,谨慎审视破坏性操作,并仅在您被允许操作所连接的硬件与网络资源的合法环境中使用本软件。
Vocat 不会绕过运营商鉴权、网络策略、硬件安全或 eSIM 信任要求。支持某项操作意味着 Vocat 能够向模组或 eUICC 发起该请求;但设备、配置文件、网络或运营商仍可能拒绝。
## 贡献
欢迎提交 Issue 与 Pull Request。请保持改动聚焦,在可行处附带测试,避免提交凭据或用户数据,并清晰地说明硬件相关行为。
提交改动前:
```bash
go test ./...
cd web && npm run build
```
## 致谢
- [Nodeseek.com](https://www.nodeseek.com) — 专注服务器的社群
- [Linux.do](https://linux.do) — 富有启发的技术社群
- [iniwex5](https://github.com/iniwex5) — 风格与功能指南
## 许可证
参见 [LICENSE](../LICENSE)。
+1
View File
@@ -3,6 +3,7 @@ module vocat
go 1.25.0
require (
github.com/coder/websocket v1.8.15
go.bug.st/serial v1.6.4
golang.org/x/crypto v0.41.0
golang.org/x/sys v0.47.0
+2
View File
@@ -1,3 +1,5 @@
github.com/coder/websocket v1.8.15 h1:6B2JPeOGlpff2Uz6vOEH1Vzpi0iUz20A+lPVhPHtNUA=
github.com/coder/websocket v1.8.15/go.mod h1:NX3SzP+inril6yawo5CQXx8+fk145lPDC6pumgx0mVg=
github.com/creack/goselect v0.1.2 h1:2DNy14+JPjRBgPzAd1thbQp4BSIihxcBf0IXhQXDRa0=
github.com/creack/goselect v0.1.2/go.mod h1:a/NhLweNvqIYMuxcMOuWY516Cimucms3DglDzQP3hKY=
github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c=
+107
View File
@@ -0,0 +1,107 @@
package developer
import (
"context"
"encoding/json"
"errors"
"fmt"
"vocat/internal/exportproxy"
"vocat/internal/httpsmode"
"vocat/internal/store"
)
func Enabled(ctx context.Context, database *store.Store) bool {
setting, err := database.AppSetting(ctx, EnabledSettingKey)
if err != nil {
return false
}
var document struct {
Enabled bool `json:"enabled"`
}
return json.Unmarshal(setting.Value, &document) == nil && document.Enabled
}
const (
EnabledSettingKey = "developer.enabled"
DeviceLimitSettingKey = "developer.device_limit"
DefaultDeviceLimit = 5
MaxDeviceLimit = 128
)
func DeviceLimit(ctx context.Context, database *store.Store, enabled bool) int {
if !enabled {
return DefaultDeviceLimit
}
setting, err := database.AppSetting(ctx, DeviceLimitSettingKey)
if err != nil {
return DefaultDeviceLimit
}
var document struct {
Limit int `json:"limit"`
}
if json.Unmarshal(setting.Value, &document) != nil || document.Limit < 1 || document.Limit > MaxDeviceLimit {
return DefaultDeviceLimit
}
return document.Limit
}
func SetDeviceLimit(ctx context.Context, database *store.Store, limit int) error {
if limit < 1 || limit > MaxDeviceLimit {
return fmt.Errorf("device limit must be between 1 and %d", MaxDeviceLimit)
}
value, err := json.Marshal(map[string]int{"limit": limit})
if err != nil {
return err
}
return database.UpsertAppSetting(ctx, store.AppSetting{Key: DeviceLimitSettingKey, Value: value})
}
// ResetExperimental restores every mutable developer-only setting. It is
// called both by `vocat develop off` and at startup whenever developer mode is
// disabled, so stale database values cannot silently remain active.
func ResetExperimental(ctx context.Context, database *store.Store) error {
httpsValue, err := json.Marshal(map[string]bool{"enabled": false})
if err != nil {
return err
}
var resetErrors []error
if err := database.UpsertAppSetting(ctx, store.AppSetting{Key: httpsmode.SettingKey, Value: httpsValue}); err != nil {
resetErrors = append(resetErrors, fmt.Errorf("reset self-signed HTTPS: %w", err))
}
if err := SetDeviceLimit(ctx, database, DefaultDeviceLimit); err != nil {
resetErrors = append(resetErrors, fmt.Errorf("reset device limit: %w", err))
}
if err := database.DeleteAppSetting(ctx, exportproxy.SettingKey); err != nil && !errors.Is(err, store.ErrNotFound) {
resetErrors = append(resetErrors, fmt.Errorf("delete export proxy configurations: %w", err))
}
devices, err := database.ListDevices(ctx)
if err != nil {
resetErrors = append(resetErrors, fmt.Errorf("list devices while disabling roaming data: %w", err))
} else {
for _, device := range devices {
if !device.NetworkEnabled {
continue
}
device.NetworkEnabled = false
if err := database.UpsertDevice(ctx, device); err != nil {
resetErrors = append(resetErrors, fmt.Errorf("disable roaming data for device %s: %w", device.ID, err))
}
}
}
policies, err := database.ListCardPolicies(ctx)
if err != nil {
resetErrors = append(resetErrors, fmt.Errorf("list card policies while disabling roaming data: %w", err))
} else {
for _, policy := range policies {
if !policy.NetworkEnabled {
continue
}
policy.NetworkEnabled = false
if err := database.UpsertCardPolicy(ctx, policy); err != nil {
resetErrors = append(resetErrors, fmt.Errorf("disable roaming policy for card %s: %w", policy.ICCID, err))
}
}
}
return errors.Join(resetErrors...)
}
+77
View File
@@ -0,0 +1,77 @@
package developer
import (
"context"
"encoding/json"
"errors"
"path/filepath"
"testing"
"vocat/internal/exportproxy"
"vocat/internal/httpsmode"
"vocat/internal/store"
)
func TestResetExperimentalRestoresDefaults(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, filepath.Join(t.TempDir(), "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := SetDeviceLimit(ctx, database, 24); err != nil {
t.Fatal(err)
}
enabled, _ := json.Marshal(map[string]bool{"enabled": true})
if err := database.UpsertAppSetting(ctx, store.AppSetting{Key: httpsmode.SettingKey, Value: enabled}); err != nil {
t.Fatal(err)
}
if err := database.UpsertDevice(ctx, store.Device{ID: "modem-1", Name: "modem-1", NetworkEnabled: true}); err != nil {
t.Fatal(err)
}
if err := database.UpsertCardPolicy(ctx, store.CardPolicy{ICCID: "8901000000000000001", NetworkEnabled: true, IPVersion: "IPV4V6"}); err != nil {
t.Fatal(err)
}
if err := database.UpsertAppSetting(ctx, store.AppSetting{Key: exportproxy.SettingKey, Value: json.RawMessage(`[]`)}); err != nil {
t.Fatal(err)
}
if err := ResetExperimental(ctx, database); err != nil {
t.Fatal(err)
}
if limit := DeviceLimit(ctx, database, true); limit != DefaultDeviceLimit {
t.Fatalf("device limit = %d, want %d", limit, DefaultDeviceLimit)
}
setting, err := database.AppSetting(ctx, httpsmode.SettingKey)
if err != nil {
t.Fatal(err)
}
var document struct {
Enabled bool `json:"enabled"`
}
if err := json.Unmarshal(setting.Value, &document); err != nil || document.Enabled {
t.Fatalf("HTTPS setting = %s, error = %v", setting.Value, err)
}
device, err := database.Device(ctx, "modem-1")
if err != nil || device.NetworkEnabled {
t.Fatalf("device roaming data was not disabled: %+v, %v", device, err)
}
policy, err := database.CardPolicy(ctx, "8901000000000000001")
if err != nil || policy.NetworkEnabled {
t.Fatalf("card roaming policy was not disabled: %+v, %v", policy, err)
}
if _, err := database.AppSetting(ctx, exportproxy.SettingKey); !errors.Is(err, store.ErrNotFound) {
t.Fatalf("export proxy configurations were not deleted: %v", err)
}
}
func TestSetDeviceLimitValidatesRange(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, filepath.Join(t.TempDir(), "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if SetDeviceLimit(ctx, database, 0) == nil || SetDeviceLimit(ctx, database, MaxDeviceLimit+1) == nil {
t.Fatal("out-of-range device limit was accepted")
}
}
+44
View File
@@ -0,0 +1,44 @@
package device
import (
_ "embed"
"encoding/json"
"strings"
)
// The offline table is generated by scripts/update-carriers.py from Android's
// versioned carrier ID database, with the previous global table retained as a
// fallback for PLMNs that Android does not yet catalogue.
//
//go:embed mccmnc.json
var carrierDatabaseJSON []byte
type carrierDatabase struct {
Carriers map[string][]string `json:"c"`
}
var globalCarrierDatabase = func() carrierDatabase {
var database carrierDatabase
if err := json.Unmarshal(carrierDatabaseJSON, &database); err != nil {
panic("device: invalid embedded MCC/MNC database: " + err.Error())
}
return database
}()
// CarrierForPLMN returns the offline carrier display name and ISO alpha-2
// country/territory code for a numeric five- or six-digit PLMN.
func CarrierForPLMN(plmn string) (name, countryCode string, ok bool) {
plmn = strings.TrimSpace(plmn)
if !decimalDigits(plmn, 5, 6) {
return "", "", false
}
entry, ok := globalCarrierDatabase.Carriers[plmn]
if !ok || len(entry) == 0 || strings.TrimSpace(entry[0]) == "" {
return "", "", false
}
name = strings.TrimSpace(entry[0])
if len(entry) > 1 {
countryCode = strings.ToUpper(strings.TrimSpace(entry[1]))
}
return name, countryCode, true
}
+181 -22
View File
@@ -23,7 +23,7 @@ func (manager *Manager) SetNetwork(
return NetworkResult{}, err
}
apn := strings.TrimSpace(request.APN)
if request.Enabled && !apnPattern.MatchString(apn) {
if request.Enabled && apn != "" && !apnPattern.MatchString(apn) {
return NetworkResult{}, ErrInvalidNetworkAPN
}
ipVersion := normalizeIPVersion(request.IPVersion)
@@ -225,7 +225,7 @@ func (manager *Manager) SetOperatorSelection(
accessTechnologyValue *int,
) (OperatorSelection, error) {
result := OperatorSelection{Mode: 0}
command := "AT+COPS=0"
command := ""
if !automatic {
plmn = strings.TrimSpace(plmn)
if len(plmn) < 5 || len(plmn) > 6 || strings.IndexFunc(plmn, func(r rune) bool { return r < '0' || r > '9' }) >= 0 {
@@ -265,28 +265,187 @@ func (manager *Manager) SetOperatorSelection(
// the lock is not aborted while registration is still in progress.
lockCtx, cancel := manager.withTimeout(ctx, manager.scanTimeout)
defer cancel()
if _, err := client.Execute(lockCtx, command); err != nil {
manager.setResult(id, state, nil, errors.New("operator selection command failed"))
if automatic {
result, err = restoreAutomaticOperatorSelection(lockCtx, client)
manager.setResult(id, state, nil, err)
return result, err
}
response, err := client.Execute(lockCtx, command)
if err != nil || !response.OK() {
if err == nil {
err = &modem.CommandError{Command: response.Command, Final: response.Final, Lines: response.Lines}
}
rollbackOperatorSelection(manager, client)
wrapped := fmt.Errorf("manual operator selection failed and automatic selection was restored: %w", err)
manager.setResult(id, state, nil, wrapped)
return OperatorSelection{}, wrapped
}
actual, err := queryOperatorSelection(lockCtx, client)
if err != nil {
rollbackOperatorSelection(manager, client)
manager.setResult(id, state, nil, err)
return OperatorSelection{}, fmt.Errorf("verify manual operator selection: %w", err)
}
if actual.Mode != 1 || actual.Operator != plmn {
rollbackOperatorSelection(manager, client)
err := fmt.Errorf("network %s did not accept registration; automatic selection was restored (modem reported mode=%d operator=%q)", plmn, actual.Mode, actual.Operator)
manager.setResult(id, state, nil, err)
return OperatorSelection{}, err
}
if !automatic {
response, err := client.Execute(lockCtx, "AT+COPS?")
if err != nil {
manager.setResult(id, state, nil, err)
return OperatorSelection{}, fmt.Errorf("verify manual operator selection: %w", err)
}
actual, err := parseOperatorSelection(response)
if err != nil {
manager.setResult(id, state, nil, err)
return OperatorSelection{}, err
}
if actual.Mode != 1 || actual.Operator != plmn {
err := fmt.Errorf("network %s did not accept registration; modem reports mode=%d operator=%q", plmn, actual.Mode, actual.Operator)
manager.setResult(id, state, nil, err)
return OperatorSelection{}, err
}
result = actual
}
result = actual
manager.setResult(id, state, nil, nil)
return result, nil
}
func queryOperatorSelection(ctx context.Context, client modem.Client) (OperatorSelection, error) {
response, err := client.Execute(ctx, "AT+COPS?")
if err != nil {
return OperatorSelection{}, err
}
if !response.OK() {
return OperatorSelection{}, &modem.CommandError{Command: response.Command, Final: response.Final, Lines: response.Lines}
}
return parseOperatorSelection(response)
}
// restoreAutomaticOperatorSelection clears both a manual PLMN latch and an
// old RAT-only scan restriction. The latter is important on EC20 modules:
// COPS=0 alone can remain effectively LTE-only after an earlier lock, unlike a
// phone's normal automatic GSM/WCDMA/LTE acquisition policy.
func restoreAutomaticOperatorSelection(ctx context.Context, client modem.Client) (OperatorSelection, error) {
// Older firmware may not implement nwscanmode; COPS auto is still useful in
// that case, so this compatibility reset is best effort.
_, _ = client.Execute(ctx, `AT+QCFG="nwscanmode",0,1`)
_, _ = client.Execute(ctx, "AT+COPS=2")
response, err := client.Execute(ctx, "AT+COPS=0")
if err != nil {
return OperatorSelection{}, err
}
if !response.OK() {
return OperatorSelection{}, &modem.CommandError{Command: response.Command, Final: response.Final, Lines: response.Lines}
}
actual, err := queryOperatorSelection(ctx, client)
if err != nil {
return OperatorSelection{}, fmt.Errorf("verify automatic operator selection: %w", err)
}
if actual.Mode != 0 {
return OperatorSelection{}, fmt.Errorf("modem did not enter automatic operator selection (mode=%d operator=%q)", actual.Mode, actual.Operator)
}
return actual, nil
}
func rollbackOperatorSelection(manager *Manager, client modem.Client) {
rollbackCtx, cancel := context.WithTimeout(context.Background(), manager.longTimeout)
defer cancel()
_, _ = restoreAutomaticOperatorSelection(rollbackCtx, client)
}
// ReRegisterOperator detaches from the network and reapplies the modem's
// current automatic/manual selection. This is intentionally different from a
// passive refresh: it forces a new registration attempt without changing the
// user's lock policy.
func (manager *Manager) ReRegisterOperator(ctx context.Context, id string) (OperatorSelection, error) {
state, err := manager.lookup(id)
if err != nil {
return OperatorSelection{}, err
}
state.opMu.Lock()
defer state.opMu.Unlock()
if err := manager.validateActive(id, state); err != nil {
return OperatorSelection{}, err
}
client, err := manager.clientLocked(ctx, state, manager.candidateFor(state))
if err != nil {
manager.setResult(id, state, nil, err)
return OperatorSelection{}, err
}
longCtx, cancel := manager.withTimeout(ctx, manager.scanTimeout)
defer cancel()
current, err := queryOperatorSelection(longCtx, client)
if err != nil {
manager.setResult(id, state, nil, err)
return OperatorSelection{}, err
}
manual := current.Mode == 1 || current.Mode == 4
if manual && !decimalPLMN(current.Operator) {
response, formatErr := client.Execute(longCtx, "AT+COPS=3,2")
if formatErr != nil || !response.OK() {
if formatErr == nil {
formatErr = &modem.CommandError{Command: response.Command, Final: response.Final, Lines: response.Lines}
}
manager.setResult(id, state, nil, formatErr)
return OperatorSelection{}, formatErr
}
current, err = queryOperatorSelection(longCtx, client)
if err != nil {
manager.setResult(id, state, nil, err)
return OperatorSelection{}, err
}
manual = current.Mode == 1 || current.Mode == 4
}
if !manual {
result, restoreErr := restoreAutomaticOperatorSelection(longCtx, client)
manager.setResult(id, state, nil, restoreErr)
return result, restoreErr
}
desired := ""
if manual {
if !decimalPLMN(current.Operator) {
return OperatorSelection{}, errors.New("current manual operator is not available as a numeric PLMN")
}
desired = fmt.Sprintf(`AT+COPS=1,2,"%s"`, current.Operator)
if code, ok := accessTechnologyCode(current.AccessTechnology); ok {
desired += fmt.Sprintf(",%d", code)
}
}
for _, command := range []string{"AT+COPS=2", desired} {
response, executeErr := client.Execute(longCtx, command)
if executeErr != nil {
manager.setResult(id, state, nil, executeErr)
return OperatorSelection{}, executeErr
}
if !response.OK() {
executeErr = &modem.CommandError{Command: response.Command, Final: response.Final, Lines: response.Lines}
manager.setResult(id, state, nil, executeErr)
return OperatorSelection{}, executeErr
}
}
result, err := queryOperatorSelection(longCtx, client)
manager.setResult(id, state, nil, err)
if err != nil {
return OperatorSelection{}, err
}
return result, nil
}
func decimalPLMN(value string) bool {
value = strings.TrimSpace(value)
return (len(value) == 5 || len(value) == 6) && strings.IndexFunc(value, func(r rune) bool {
return r < '0' || r > '9'
}) < 0
}
func accessTechnologyCode(name string) (int, bool) {
switch strings.ToUpper(strings.TrimSpace(name)) {
case "GSM":
return 0, true
case "UTRAN":
return 2, true
case "EDGE":
return 3, true
case "HSDPA":
return 4, true
case "HSUPA":
return 5, true
case "HSPA":
return 6, true
case "LTE":
return 7, true
case "NR5G":
return 9, true
default:
return 0, false
}
}
+208 -25
View File
@@ -4,9 +4,13 @@ package device
import (
"context"
"errors"
"fmt"
"hash/fnv"
"net"
"os"
"os/exec"
"strconv"
"strings"
"time"
@@ -31,7 +35,11 @@ func setQMINetwork(
profilePath := profile.Name()
defer os.Remove(profilePath)
ipType := map[string]string{"IP": "4", "IPV6": "6", "IPV4V6": "4"}[ipVersion]
if _, err := fmt.Fprintf(profile, "APN=%s\nIP_TYPE=%s\nPROXY=yes\n", apn, ipType); err != nil {
profileText := fmt.Sprintf("IP_TYPE=%s\nPROXY=yes\n", ipType)
if apn != "" {
profileText = "APN=" + apn + "\n" + profileText
}
if _, err := fmt.Fprint(profile, profileText); err != nil {
_ = profile.Close()
return NetworkResult{}, fmt.Errorf("write temporary QMI profile: %w", err)
}
@@ -54,36 +62,42 @@ func setQMINetwork(
lowerDetail := strings.ToLower(detail)
idempotentStop := !enabled && (strings.Contains(lowerDetail, "already stopped") ||
strings.Contains(lowerDetail, "not started") || strings.Contains(lowerDetail, "no network"))
if !idempotentStop {
idempotentStart := enabled && (strings.Contains(lowerDetail, "already started") ||
strings.Contains(lowerDetail, "already connected"))
if !idempotentStop && !idempotentStart {
return NetworkResult{}, fmt.Errorf("qmi-network %s failed: %w: %s", action, err, detail)
}
}
if ipCommand, lookErr := exec.LookPath("ip"); lookErr == nil {
linkAction := "down"
if enabled {
linkAction = "up"
}
linkOutput, linkErr := exec.CommandContext(ctx, ipCommand, "link", "set", "dev", candidate.NetworkInterface, linkAction).CombinedOutput()
if linkErr != nil {
return NetworkResult{}, fmt.Errorf("set %s %s: %w: %s", candidate.NetworkInterface, linkAction, linkErr, strings.TrimSpace(string(linkOutput)))
}
ipCommand, lookErr := exec.LookPath("ip")
if lookErr != nil {
return NetworkResult{}, fmt.Errorf("%w: install iproute2 to control %s", ErrDataBackendUnavailable, candidate.NetworkInterface)
}
linkAction := "down"
if enabled {
linkAction = "up"
}
linkOutput, linkErr := exec.CommandContext(ctx, ipCommand, "link", "set", "dev", candidate.NetworkInterface, linkAction).CombinedOutput()
if linkErr != nil {
return NetworkResult{}, fmt.Errorf("set %s %s: %w: %s", candidate.NetworkInterface, linkAction, linkErr, strings.TrimSpace(string(linkOutput)))
}
if enabled {
if busybox, lookErr := exec.LookPath("busybox"); lookErr == nil {
dhcpOutput, dhcpErr := exec.CommandContext(ctx, busybox, "udhcpc", "-q", "-n", "-t", "5", "-T", "3", "-i", candidate.NetworkInterface).CombinedOutput()
if dhcpErr != nil {
rollbackCtx, cancelRollback := context.WithTimeout(context.Background(), managerCommandCleanupTimeout)
defer cancelRollback()
_, _ = exec.CommandContext(rollbackCtx, qmiNetwork, "--profile="+profilePath, candidate.QMIControl, "stop").CombinedOutput()
if ipCommand, lookErr := exec.LookPath("ip"); lookErr == nil {
_, _ = exec.CommandContext(rollbackCtx, ipCommand, "link", "set", "dev", candidate.NetworkInterface, "down").CombinedOutput()
}
return NetworkResult{}, fmt.Errorf("QMI session started but DHCP failed: %w: %s", dhcpErr, strings.TrimSpace(string(dhcpOutput)))
}
if value := strings.TrimSpace(string(dhcpOutput)); value != "" {
detail = strings.TrimSpace(detail + "\n" + value)
}
busybox, busyboxErr := exec.LookPath("busybox")
if busyboxErr != nil {
return NetworkResult{}, fmt.Errorf("%w: busybox udhcpc is required for %s", ErrDataBackendUnavailable, candidate.NetworkInterface)
}
dhcpDetail, dhcpErr := configureExportProxyDHCP(ctx, busybox, ipCommand, candidate.NetworkInterface)
if dhcpErr != nil {
rollbackCtx, cancelRollback := context.WithTimeout(context.Background(), managerCommandCleanupTimeout)
defer cancelRollback()
clearExportProxyRoute(rollbackCtx, candidate.NetworkInterface)
_, _ = exec.CommandContext(rollbackCtx, qmiNetwork, "--profile="+profilePath, candidate.QMIControl, "stop").CombinedOutput()
_, _ = exec.CommandContext(rollbackCtx, ipCommand, "link", "set", "dev", candidate.NetworkInterface, "down").CombinedOutput()
return NetworkResult{}, fmt.Errorf("QMI session started but protected DHCP failed: %w", dhcpErr)
}
detail = strings.TrimSpace(detail + "\n" + dhcpDetail)
} else {
clearExportProxyRoute(ctx, candidate.NetworkInterface)
_, _ = exec.CommandContext(ctx, ipCommand, "-4", "addr", "flush", "dev", candidate.NetworkInterface, "scope", "global").CombinedOutput()
}
return NetworkResult{
Enabled: enabled,
@@ -96,4 +110,173 @@ func setQMINetwork(
}, nil
}
// exportProxyRouteIdentity must stay in sync with the Export Proxy plugin's
// Linux socket mark. Unmarked host traffic never sees the cellular default
// route; only plugin sockets carrying this mark are policy-routed to it.
func exportProxyRouteIdentity(networkInterface string) (mark uint32, table, priority int) {
hash := fnv.New32a()
_, _ = hash.Write([]byte(networkInterface))
value := hash.Sum32()
mark = 0x56000000 | (value & 0x00ffffff)
table = 20000 + int(value%10000)
priority = 20000 + int(value%10000)
return
}
func configureExportProxyDHCP(ctx context.Context, busybox, ipCommand, networkInterface string) (string, error) {
lease, err := os.CreateTemp("", "vocat-dhcp-lease-*.env")
if err != nil {
return "", err
}
leasePath := lease.Name()
_ = lease.Close()
_ = os.Remove(leasePath)
defer os.Remove(leasePath)
script, err := os.CreateTemp("", "vocat-udhcpc-*.sh")
if err != nil {
return "", err
}
scriptPath := script.Name()
defer os.Remove(scriptPath)
scriptText := fmt.Sprintf(`#!/bin/sh
case "$1" in
bound|renew)
(umask 077; printf 'ip=%%s\nsubnet=%%s\nrouter=%%s\ndns=%%s\n' "$ip" "$subnet" "$router" "$dns" > %q)
;;
esac
exit 0
`, leasePath)
if _, err := script.WriteString(scriptText); err != nil {
_ = script.Close()
return "", err
}
if err := script.Chmod(0o700); err != nil {
_ = script.Close()
return "", err
}
if err := script.Close(); err != nil {
return "", err
}
output, err := exec.CommandContext(ctx, busybox, "udhcpc", "-q", "-n", "-t", "5", "-T", "3", "-i", networkInterface, "-s", scriptPath).CombinedOutput()
if err != nil {
if strings.Contains(strings.ToLower(string(output)), "address family not supported") {
return "", fmt.Errorf("udhcpc cannot open its link-layer socket: allow AF_PACKET in the vocat systemd service RestrictAddressFamilies setting: %w", err)
}
return "", fmt.Errorf("udhcpc: %w: %s", err, strings.TrimSpace(string(output)))
}
raw, err := os.ReadFile(leasePath)
if err != nil {
return "", fmt.Errorf("read DHCP lease: %w", err)
}
values := make(map[string]string)
for _, line := range strings.Split(string(raw), "\n") {
key, value, found := strings.Cut(line, "=")
if found {
values[strings.TrimSpace(key)] = strings.TrimSpace(value)
}
}
address := net.ParseIP(values["ip"]).To4()
maskIP := net.ParseIP(values["subnet"]).To4()
if address == nil || maskIP == nil {
return "", errors.New("DHCP returned no valid IPv4 address/subnet")
}
mask := net.IPMask(maskIP)
ones, bits := mask.Size()
if bits != 32 || ones < 0 {
return "", errors.New("DHCP returned an invalid IPv4 subnet")
}
network := address.Mask(mask)
routers := strings.Fields(values["router"])
if len(routers) > 0 && net.ParseIP(routers[0]).To4() == nil {
return "", errors.New("DHCP returned an invalid IPv4 gateway")
}
if result, addrErr := exec.CommandContext(ctx, ipCommand, "-4", "addr", "replace", fmt.Sprintf("%s/%d", address.String(), ones), "dev", networkInterface).CombinedOutput(); addrErr != nil {
return "", fmt.Errorf("configure cellular address: %w: %s", addrErr, strings.TrimSpace(string(result)))
}
mark, table, priority := exportProxyRouteIdentity(networkInterface)
clearExportProxyRoute(ctx, networkInterface)
connectedCIDR := fmt.Sprintf("%s/%d", network.String(), ones)
if result, routeErr := exec.CommandContext(ctx, ipCommand, "-4", "route", "replace", "table", strconv.Itoa(table), connectedCIDR, "dev", networkInterface, "scope", "link", "src", address.String()).CombinedOutput(); routeErr != nil {
clearExportProxyRoute(ctx, networkInterface)
return "", fmt.Errorf("install protected connected route: %w: %s", routeErr, strings.TrimSpace(string(result)))
}
defaultArgs := []string{"-4", "route", "replace", "table", strconv.Itoa(table), "default"}
if len(routers) > 0 {
defaultArgs = append(defaultArgs, "via", routers[0])
}
defaultArgs = append(defaultArgs, "dev", networkInterface, "onlink")
if result, routeErr := exec.CommandContext(ctx, ipCommand, defaultArgs...).CombinedOutput(); routeErr != nil {
clearExportProxyRoute(ctx, networkInterface)
return "", fmt.Errorf("install protected default route: %w: %s", routeErr, strings.TrimSpace(string(result)))
}
markText := fmt.Sprintf("0x%x", mark)
result, err := exec.CommandContext(ctx, ipCommand, "rule", "add", "priority", strconv.Itoa(priority), "fwmark", markText, "lookup", strconv.Itoa(table)).CombinedOutput()
if err != nil {
clearExportProxyRoute(ctx, networkInterface)
return "", fmt.Errorf("install protected routing rule: %w: %s", err, strings.TrimSpace(string(result)))
}
if err := writeExportProxyDNS(networkInterface, strings.Fields(values["dns"])); err != nil {
clearExportProxyRoute(ctx, networkInterface)
return "", fmt.Errorf("publish protected DNS configuration: %w", err)
}
return fmt.Sprintf("protected DHCP lease %s/%d", address.String(), ones), nil
}
func exportProxyDNSPath(networkInterface string) string {
safeName := strings.Map(func(character rune) rune {
if character >= 'a' && character <= 'z' || character >= 'A' && character <= 'Z' ||
character >= '0' && character <= '9' || character == '-' || character == '_' || character == '.' {
return character
}
return '_'
}, networkInterface)
return "/run/vocat/cellular-" + safeName + ".dns"
}
func writeExportProxyDNS(networkInterface string, servers []string) error {
valid := make([]string, 0, len(servers))
for _, server := range servers {
if address := net.ParseIP(server); address != nil {
valid = append(valid, address.String())
}
}
if len(valid) == 0 {
// This is used only by marked Export Proxy sockets. It never changes the
// host resolver and is merely a fallback for carriers omitting DHCP DNS.
valid = []string{"1.1.1.1", "8.8.8.8"}
}
if err := os.MkdirAll("/run/vocat", 0o755); err != nil {
return err
}
temporary, err := os.CreateTemp("/run/vocat", ".cellular-dns-*")
if err != nil {
return err
}
temporaryPath := temporary.Name()
defer os.Remove(temporaryPath)
if _, err := temporary.WriteString(strings.Join(valid, "\n") + "\n"); err != nil {
_ = temporary.Close()
return err
}
if err := temporary.Chmod(0o644); err != nil {
_ = temporary.Close()
return err
}
if err := temporary.Close(); err != nil {
return err
}
return os.Rename(temporaryPath, exportProxyDNSPath(networkInterface))
}
func clearExportProxyRoute(ctx context.Context, networkInterface string) {
_ = os.Remove(exportProxyDNSPath(networkInterface))
ipCommand, err := exec.LookPath("ip")
if err != nil {
return
}
mark, table, priority := exportProxyRouteIdentity(networkInterface)
_, _ = exec.CommandContext(ctx, ipCommand, "rule", "del", "priority", strconv.Itoa(priority), "fwmark", fmt.Sprintf("0x%x", mark), "lookup", strconv.Itoa(table)).CombinedOutput()
_, _ = exec.CommandContext(ctx, ipCommand, "-4", "route", "flush", "table", strconv.Itoa(table)).CombinedOutput()
}
const managerCommandCleanupTimeout = 15 * time.Second
+82 -1
View File
@@ -74,7 +74,10 @@ func TestOperatorSelectionManualAndAutomatic(t *testing.T) {
client := &transcriptClient{steps: []clientStep{
{command: `AT+COPS=1,2,"46000",7`, response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 1,2,"46000",7`)},
{command: `AT+QCFG="nwscanmode",0,1`, response: okResponse()},
{command: "AT+COPS=2", response: okResponse()},
{command: "AT+COPS=0", response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 0,2,"46001",7`)},
}}
manager, id := newStartedTestManager(t, client)
act := 7
@@ -89,7 +92,7 @@ func TestOperatorSelectionManualAndAutomatic(t *testing.T) {
if err != nil {
t.Fatalf("automatic selection: %v", err)
}
if selection.Mode != 0 || selection.Operator != "" {
if selection.Mode != 0 || selection.Operator != "46001" {
t.Fatalf("automatic selection = %#v", selection)
}
client.assertDone(t)
@@ -99,6 +102,10 @@ func TestOperatorSelectionRejectsAutomaticFallbackAsSuccess(t *testing.T) {
client := &transcriptClient{steps: []clientStep{
{command: `AT+COPS=1,2,"46000",7`, response: okResponse()},
{command: "AT+COPS?", response: okResponse("+COPS: 0")},
{command: `AT+QCFG="nwscanmode",0,1`, response: okResponse()},
{command: "AT+COPS=2", response: okResponse()},
{command: "AT+COPS=0", response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 0,2,"46001",7`)},
}}
manager, id := newStartedTestManager(t, client)
act := 7
@@ -107,3 +114,77 @@ func TestOperatorSelectionRejectsAutomaticFallbackAsSuccess(t *testing.T) {
}
client.assertDone(t)
}
func TestOperatorSelectionCommandFailureRestoresAutomaticMode(t *testing.T) {
selectionErr := errors.New("+CME ERROR: 30")
client := &transcriptClient{steps: []clientStep{
{command: `AT+COPS=1,2,"46000",7`, err: selectionErr},
{command: `AT+QCFG="nwscanmode",0,1`, response: okResponse()},
{command: "AT+COPS=2", response: okResponse()},
{command: "AT+COPS=0", response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 0,2,"46001",7`)},
}}
manager, id := newStartedTestManager(t, client)
act := 7
_, err := manager.SetOperatorSelection(context.Background(), id, false, "46000", &act)
if !errors.Is(err, selectionErr) {
t.Fatalf("error = %v, want wrapped selection error", err)
}
client.assertDone(t)
}
func TestReRegisterOperatorReappliesAutomaticMode(t *testing.T) {
client := &transcriptClient{steps: []clientStep{
{command: "AT+COPS?", response: okResponse(`+COPS: 0,2,"46001",7`)},
{command: `AT+QCFG="nwscanmode",0,1`, response: okResponse()},
{command: "AT+COPS=2", response: okResponse()},
{command: "AT+COPS=0", response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 0,2,"46001",7`)},
}}
manager, id := newStartedTestManager(t, client)
selection, err := manager.ReRegisterOperator(context.Background(), id)
if err != nil {
t.Fatal(err)
}
if selection.Mode != 0 || selection.Operator != "46001" {
t.Fatalf("selection = %#v", selection)
}
client.assertDone(t)
}
func TestReRegisterOperatorPreservesManualLock(t *testing.T) {
client := &transcriptClient{steps: []clientStep{
{command: "AT+COPS?", response: okResponse(`+COPS: 1,2,"46003",7`)},
{command: "AT+COPS=2", response: okResponse()},
{command: `AT+COPS=1,2,"46003",7`, response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 1,2,"46003",7`)},
}}
manager, id := newStartedTestManager(t, client)
selection, err := manager.ReRegisterOperator(context.Background(), id)
if err != nil {
t.Fatal(err)
}
if selection.Mode != 1 || selection.Operator != "46003" || selection.AccessTechnology != "LTE" {
t.Fatalf("selection = %#v", selection)
}
client.assertDone(t)
}
func TestReRegisterOperatorRecoversDeregisteredModeWithAutomaticSelection(t *testing.T) {
client := &transcriptClient{steps: []clientStep{
{command: "AT+COPS?", response: okResponse(`+COPS: 2`)},
{command: `AT+QCFG="nwscanmode",0,1`, response: okResponse()},
{command: "AT+COPS=2", response: okResponse()},
{command: "AT+COPS=0", response: okResponse()},
{command: "AT+COPS?", response: okResponse(`+COPS: 0,2,"46001",7`)},
}}
manager, id := newStartedTestManager(t, client)
selection, err := manager.ReRegisterOperator(context.Background(), id)
if err != nil {
t.Fatal(err)
}
if selection.Mode != 0 || selection.Operator != "46001" {
t.Fatalf("selection = %#v", selection)
}
client.assertDone(t)
}
+4 -1
View File
@@ -27,6 +27,7 @@ func TestManagerRefreshBuildsEC20Snapshot(t *testing.T) {
),
},
{command: "AT+COPS?", response: okResponse(`+COPS: 0,0,"China Mobile",7`)},
{command: "AT+CEREG?", response: okResponse(`+CEREG: 0,5`)},
{command: "AT+CGSN", response: okResponse("867123456789012")},
{
command: "AT+CCID",
@@ -66,7 +67,9 @@ func TestManagerRefreshBuildsEC20Snapshot(t *testing.T) {
t.Fatalf("signal metrics = %#v", snapshot)
}
if snapshot.AccessTech != "LTE" || snapshot.Band != "B3" ||
snapshot.Channel != "1650" || snapshot.OperatorName != "China Mobile" {
snapshot.Channel != "1650" || snapshot.OperatorName != "China Unicom" ||
snapshot.OperatorCode != "46001" ||
snapshot.RegistrationStatus != 5 || snapshot.RegistrationSource != "CEREG" {
t.Fatalf("network = %#v", snapshot)
}
if snapshot.IMEI != "867123456789012" ||
File diff suppressed because one or more lines are too long
+36
View File
@@ -0,0 +1,36 @@
//go:build linux
package device
import (
"context"
"os/exec"
"strings"
"time"
"vocat/internal/modem"
)
func readPlatformRegistration(ctx context.Context, candidate modem.Candidate) (platformRegistration, bool) {
control := strings.TrimSpace(candidate.QMIControl)
if control == "" {
return platformRegistration{}, false
}
qmicli, err := exec.LookPath("qmicli")
if err != nil {
return platformRegistration{}, false
}
queryContext, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
output, err := exec.CommandContext(
queryContext,
qmicli,
"-d", control,
"--device-open-proxy",
"--nas-get-serving-system",
).CombinedOutput()
if err != nil {
return platformRegistration{}, false
}
return parseQMIRegistration(string(output))
}
+13
View File
@@ -0,0 +1,13 @@
//go:build !linux
package device
import (
"context"
"vocat/internal/modem"
)
func readPlatformRegistration(context.Context, modem.Candidate) (platformRegistration, bool) {
return platformRegistration{}, false
}
+80
View File
@@ -0,0 +1,80 @@
package device
import (
"regexp"
"strings"
)
type platformRegistration struct {
Status int
PLMN string
Name string
PSAttached bool
}
var qmiQuotedFieldPattern = regexp.MustCompile(`(?i)^\s*([^:]+):\s*'([^']*)'\s*$`)
func parseQMIRegistration(output string) (platformRegistration, bool) {
result := platformRegistration{}
registrationState := ""
roaming := false
mcc := ""
mnc := ""
pcsDigit := false
for _, rawLine := range strings.Split(output, "\n") {
match := qmiQuotedFieldPattern.FindStringSubmatch(strings.TrimSpace(rawLine))
if len(match) != 3 {
continue
}
key := strings.ToLower(strings.TrimSpace(match[1]))
value := strings.TrimSpace(match[2])
switch key {
case "registration state":
registrationState = strings.ToLower(value)
case "roaming status":
roaming = strings.EqualFold(value, "on")
case "ps":
result.PSAttached = strings.EqualFold(value, "attached")
case "mcc":
if mcc == "" {
mcc = value
}
case "mnc":
if mnc == "" {
mnc = value
}
case "description":
if result.Name == "" {
result.Name = value
}
case "mnc with pcs digit":
pcsDigit = strings.EqualFold(value, "yes")
}
}
switch registrationState {
case "registered":
result.Status = 1
if roaming {
result.Status = 5
}
case "not-registered-searching", "searching":
result.Status = 2
case "registration-denied", "denied":
result.Status = 3
case "not-registered":
result.Status = 0
default:
return platformRegistration{}, false
}
if decimalDigits(mcc, 3, 3) && decimalDigits(mnc, 1, 3) {
width := 2
if pcsDigit {
width = 3
}
for len(mnc) < width {
mnc = "0" + mnc
}
result.PLMN = mcc + mnc
}
return result, true
}
+31
View File
@@ -0,0 +1,31 @@
package device
import "testing"
func TestParseQMIRegistrationRegisteredRoaming(t *testing.T) {
output := `
Registration state: 'registered'
CS: 'attached'
PS: 'attached'
Roaming status: 'on'
Current PLMN:
MCC: '460'
MNC: '1'
Description: 'UNICOM'
Full operator code info:
MCC: '460'
MNC: '1'
MNC with PCS digit: 'no'
`
result, found := parseQMIRegistration(output)
if !found || result.Status != 5 || !result.PSAttached || result.PLMN != "46001" || result.Name != "UNICOM" {
t.Fatalf("registration = %#v, found=%v", result, found)
}
}
func TestParseQMIRegistrationSearching(t *testing.T) {
result, found := parseQMIRegistration("Registration state: 'not-registered-searching'\nPS: 'detached'")
if !found || result.Status != 2 || result.PSAttached {
t.Fatalf("registration = %#v, found=%v", result, found)
}
}
+24
View File
@@ -0,0 +1,24 @@
package device
import (
"testing"
"vocat/internal/modem"
)
func TestParseRegistrationStatus(t *testing.T) {
tests := []struct {
line string
want int
}{
{line: "+CEREG: 0,5", want: 5},
{line: "+CGREG: 2,1,\"FFFE\",\"06698D06\",7", want: 1},
{line: "+CREG: 2", want: 2},
}
for _, test := range tests {
got, ok := parseRegistrationStatus(modem.Response{Lines: []string{test.line}})
if !ok || got != test.want {
t.Fatalf("parseRegistrationStatus(%q) = %d, %v", test.line, got, ok)
}
}
}
+24 -1
View File
@@ -14,6 +14,7 @@ type ScannedOperator struct {
Name string `json:"name"`
Short string `json:"shortName,omitempty"`
Numeric string `json:"numeric"`
Country string `json:"countryCode,omitempty"`
Act string `json:"act,omitempty"`
}
@@ -75,11 +76,19 @@ func parseOperatorScan(response modem.Response) []ScannedOperator {
if len(fields) < 4 {
continue
}
name, country, _ := CarrierForPLMN(fields[3])
if name == "" {
name = strings.TrimSpace(fields[1])
}
if name == "" {
name = strings.TrimSpace(fields[3])
}
operator := ScannedOperator{
Status: operatorScanStatus(fields[0]),
Name: fields[1],
Name: name,
Short: fields[2],
Numeric: fields[3],
Country: country,
}
if len(fields) >= 5 {
operator.Act = accessTechnology(fields[4])
@@ -90,6 +99,20 @@ func parseOperatorScan(response modem.Response) []ScannedOperator {
return operators
}
// carrierNameForPLMN resolves the numeric serving PLMN through the bundled
// global carrier database. Some EC20 firmware returns an empty, localized, or
// stale long name even though the MCC/MNC is correct. The numeric identity is
// the authoritative value used for network selection.
func carrierNameForPLMN(plmn, fallback string) string {
if name, _, ok := CarrierForPLMN(plmn); ok {
return name
}
if fallback = strings.TrimSpace(fallback); fallback != "" {
return fallback
}
return strings.TrimSpace(plmn)
}
// extractScanTuples returns the contents of each top-level parenthesised group,
// ignoring parentheses inside quoted strings.
func extractScanTuples(payload string) []string {
+55
View File
@@ -0,0 +1,55 @@
package device
import (
"testing"
"vocat/internal/modem"
)
func TestParseOperatorScanNormalizesMainlandCarrierNamesByPLMN(t *testing.T) {
response := modem.Response{Lines: []string{
`+COPS: (1,"CMCC","CMCC","46000",7),(1,"wrong modem name","CU","46001",7),(1,"","CT","46011",7),(1,"CBN","CBN","46015",7)`,
}}
operators := parseOperatorScan(response)
if len(operators) != 4 {
t.Fatalf("operators = %#v", operators)
}
want := []string{"China Mobile", "China Unicom", "China Telecom", "China Broadnet"}
for index := range want {
if operators[index].Name != want[index] {
t.Fatalf("operator %d name = %q, want %q", index, operators[index].Name, want[index])
}
}
}
func TestCarrierNameForPLMNUsesGlobalDatabase(t *testing.T) {
if got := carrierNameForPLMN("23415", "stale modem name"); got != "Vodafone" {
t.Fatalf("carrier name = %q", got)
}
if got := carrierNameForPLMN("26202", ""); got != "Vodafone" {
t.Fatalf("German carrier name = %q", got)
}
if got := carrierNameForPLMN("310260", ""); got != "T-Mobile - US" {
t.Fatalf("US carrier name = %q", got)
}
if got := carrierNameForPLMN("99999", "Test Network"); got != "Test Network" {
t.Fatalf("unknown carrier fallback = %q", got)
}
}
func TestCarrierForPLMNReturnsCountryCode(t *testing.T) {
tests := map[string]string{
"23415": "GB",
"26202": "DE",
"310260": "US",
"22201": "IT",
"72405": "BR",
"46015": "CN",
}
for plmn, wantCountry := range tests {
name, country, ok := CarrierForPLMN(plmn)
if !ok || name == "" || country != wantCountry {
t.Errorf("CarrierForPLMN(%q) = (%q, %q, %v), want a name and country %q", plmn, name, country, ok, wantCountry)
}
}
}
+64 -2
View File
@@ -50,8 +50,10 @@ func (manager *Manager) readSnapshot(
if response, ok := optional("AT+CSQ"); ok {
snapshot.SignalRaw, snapshot.SignalPercent, snapshot.RSSIDBm = parseCSQ(response)
}
servingPLMN := ""
if response, ok := optional(`AT+QENG="servingcell"`); ok {
metrics := parseQENG(response)
servingPLMN = metrics.PLMN
snapshot.AccessTech = metrics.AccessTech
snapshot.Band = metrics.Band
snapshot.Channel = metrics.Channel
@@ -64,12 +66,42 @@ func (manager *Manager) readSnapshot(
}
if response, ok := optional("AT+COPS?"); ok {
operator := parseCOPS(response)
snapshot.OperatorName = operator.Name
snapshot.OperatorCode = operator.Code
if operator.Code != "" {
snapshot.OperatorCode = operator.Code
} else {
snapshot.OperatorCode = servingPLMN
}
snapshot.OperatorName = carrierNameForPLMN(snapshot.OperatorCode, operator.Name)
if snapshot.AccessTech == "" {
snapshot.AccessTech = operator.AccessTech
}
}
for _, command := range []string{"AT+CEREG?", "AT+CGREG?", "AT+CREG?"} {
response, registrationErr := manager.command(ctx, client, command)
if registrationErr != nil {
continue
}
if status, found := parseRegistrationStatus(response); found {
snapshot.RegistrationStatus = status
snapshot.RegistrationSource = strings.TrimSuffix(strings.TrimPrefix(command, "AT+"), "?")
break
}
}
if registration, found := readPlatformRegistration(ctx, candidate); found {
snapshot.RegistrationStatus = registration.Status
snapshot.RegistrationSource = "QMI NAS"
snapshot.PSAttached = registration.PSAttached
if registration.PLMN != "" {
snapshot.OperatorCode = registration.PLMN
snapshot.OperatorName = carrierNameForPLMN(registration.PLMN, registration.Name)
}
}
if snapshot.RegistrationSource == "" && (snapshot.OperatorName != "" || snapshot.OperatorCode != "") {
// Older firmware can omit registration queries while COPS still proves
// that an operator is selected.
snapshot.RegistrationStatus = 1
snapshot.RegistrationSource = "COPS"
}
if response, ok := optional("AT+CGSN"); ok {
snapshot.IMEI = parseIdentifier(
response,
@@ -107,6 +139,25 @@ func (manager *Manager) readSnapshot(
return snapshot, nil
}
func parseRegistrationStatus(response modem.Response) (int, bool) {
for _, prefix := range []string{"+CEREG:", "+CGREG:", "+CREG:"} {
values := csvValues(valueAfterPrefix(response, prefix))
if len(values) == 0 {
continue
}
index := 0
// Query responses are <n>,<stat>; unsolicited responses are <stat>.
if len(values) >= 2 {
index = 1
}
status, err := strconv.Atoi(strings.TrimSpace(values[index]))
if err == nil && status >= 0 && status <= 10 {
return status, true
}
}
return 0, false
}
func parseATI(lines []string) (manufacturer, model, firmware string) {
for _, line := range lines {
line = strings.TrimSpace(line)
@@ -159,6 +210,7 @@ func parseCSQ(response modem.Response) (raw, percent, dbm *int) {
}
type qengMetrics struct {
PLMN string
AccessTech string
Band string
Channel string
@@ -179,6 +231,9 @@ func parseQENG(response modem.Response) qengMetrics {
}
result := qengMetrics{AccessTech: strings.ToUpper(values[2])}
if strings.EqualFold(values[2], "LTE") && len(values) >= 17 {
if decimalDigits(values[4], 3, 3) && decimalDigits(values[5], 2, 3) {
result.PLMN = values[4] + values[5]
}
result.Channel = values[8]
if values[9] != "" {
result.Band = "B" + values[9]
@@ -193,6 +248,13 @@ func parseQENG(response modem.Response) qengMetrics {
return qengMetrics{}
}
func decimalDigits(value string, minimum, maximum int) bool {
value = strings.TrimSpace(value)
return len(value) >= minimum && len(value) <= maximum && strings.IndexFunc(value, func(character rune) bool {
return character < '0' || character > '9'
}) < 0
}
type operatorInfo struct {
Name string
Code string
+32 -29
View File
@@ -73,35 +73,38 @@ const (
)
type Snapshot struct {
DeviceID string `json:"deviceId"`
Port string `json:"port"`
Responsive bool `json:"responsive"`
Manufacturer string `json:"manufacturer"`
Model string `json:"model"`
Firmware string `json:"firmware"`
SIMStatus string `json:"simStatus"`
SIMReady bool `json:"simReady"`
SignalRaw *int `json:"signalRaw,omitempty"`
SignalPercent *int `json:"signalPercent,omitempty"`
RSSIDBm *int `json:"rssiDbm,omitempty"`
RSRP *int `json:"rsrp,omitempty"`
RSRQ *int `json:"rsrq,omitempty"`
SINR *int `json:"sinr,omitempty"`
AccessTech string `json:"accessTech"`
Band string `json:"band"`
Channel string `json:"channel"`
OperatorName string `json:"operatorName"`
OperatorCode string `json:"operatorCode"`
IMEI string `json:"imei"`
ICCID string `json:"iccid"`
IMSI string `json:"imsi"`
OperatingMode int `json:"operatingMode"`
ModeKnown bool `json:"modeKnown"`
FlightMode bool `json:"flightMode"`
RadioOff bool `json:"radioOff"`
Phone PhoneNumber `json:"phone"`
Warnings []string `json:"warnings,omitempty"`
UpdatedAt time.Time `json:"updatedAt"`
DeviceID string `json:"deviceId"`
Port string `json:"port"`
Responsive bool `json:"responsive"`
Manufacturer string `json:"manufacturer"`
Model string `json:"model"`
Firmware string `json:"firmware"`
SIMStatus string `json:"simStatus"`
SIMReady bool `json:"simReady"`
SignalRaw *int `json:"signalRaw,omitempty"`
SignalPercent *int `json:"signalPercent,omitempty"`
RSSIDBm *int `json:"rssiDbm,omitempty"`
RSRP *int `json:"rsrp,omitempty"`
RSRQ *int `json:"rsrq,omitempty"`
SINR *int `json:"sinr,omitempty"`
AccessTech string `json:"accessTech"`
Band string `json:"band"`
Channel string `json:"channel"`
OperatorName string `json:"operatorName"`
OperatorCode string `json:"operatorCode"`
RegistrationStatus int `json:"registrationStatus"`
RegistrationSource string `json:"registrationSource"`
PSAttached bool `json:"psAttached"`
IMEI string `json:"imei"`
ICCID string `json:"iccid"`
IMSI string `json:"imsi"`
OperatingMode int `json:"operatingMode"`
ModeKnown bool `json:"modeKnown"`
FlightMode bool `json:"flightMode"`
RadioOff bool `json:"radioOff"`
Phone PhoneNumber `json:"phone"`
Warnings []string `json:"warnings,omitempty"`
UpdatedAt time.Time `json:"updatedAt"`
}
type USSDResult struct {
+80
View File
@@ -0,0 +1,80 @@
//go:build linux
package exportproxy
import (
"bufio"
"context"
"hash/fnv"
"net"
"os"
"path/filepath"
"strings"
"syscall"
)
func platformSupported() error { return nil }
func boundDialer(networkInterface string) net.Dialer {
return net.Dialer{Control: func(_, _ string, raw syscall.RawConn) error {
var bindError error
err := raw.Control(func(fd uintptr) {
if err := syscall.SetsockoptInt(int(fd), syscall.SOL_SOCKET, syscall.SO_MARK, int(exportRouteMark(networkInterface))); err != nil {
bindError = err
return
}
bindError = syscall.SetsockoptString(int(fd), syscall.SOL_SOCKET, syscall.SO_BINDTODEVICE, networkInterface)
})
if err != nil {
return err
}
return bindError
}}
}
func exportRouteMark(networkInterface string) uint32 {
hash := fnv.New32a()
_, _ = hash.Write([]byte(networkInterface))
return 0x56000000 | (hash.Sum32() & 0x00ffffff)
}
func boundResolver(networkInterface string) *net.Resolver {
dialer := boundDialer(networkInterface)
return &net.Resolver{PreferGo: true, Dial: func(ctx context.Context, network, _ string) (net.Conn, error) {
var lastError error
for _, server := range exportRouteDNSServers(networkInterface) {
connection, err := dialer.DialContext(ctx, network, net.JoinHostPort(server, "53"))
if err == nil {
return connection, nil
}
lastError = err
}
return nil, lastError
}}
}
func exportRouteDNSServers(networkInterface string) []string {
safeName := strings.Map(func(character rune) rune {
if character >= 'a' && character <= 'z' || character >= 'A' && character <= 'Z' ||
character >= '0' && character <= '9' || character == '-' || character == '_' || character == '.' {
return character
}
return '_'
}, networkInterface)
file, err := os.Open(filepath.Join("/run/vocat", "cellular-"+safeName+".dns"))
if err != nil {
return []string{"1.1.1.1", "8.8.8.8"}
}
defer file.Close()
servers := make([]string, 0, 2)
scanner := bufio.NewScanner(file)
for scanner.Scan() {
if value := strings.TrimSpace(scanner.Text()); net.ParseIP(value) != nil {
servers = append(servers, value)
}
}
if len(servers) == 0 {
return []string{"1.1.1.1", "8.8.8.8"}
}
return servers
}
+12
View File
@@ -0,0 +1,12 @@
//go:build !linux
package exportproxy
import (
"errors"
"net"
)
func platformSupported() error { return errors.New("built-in export proxy is only available on Linux") }
func boundDialer(string) net.Dialer { return net.Dialer{} }
func boundResolver(string) *net.Resolver { return net.DefaultResolver }
+88
View File
@@ -0,0 +1,88 @@
package exportproxy
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"net/http"
"strings"
"time"
)
const ipInfoURL = "https://ipinfo.io/json"
type PublicIPInfo struct {
IP string `json:"ip"`
CountryCode string `json:"country_code"`
Region string `json:"region"`
City string `json:"city"`
Organization string `json:"organization,omitempty"`
}
// LookupPublicIP sends the lookup through the same marked, interface-bound
// dialer and isolated DNS resolver as Export Proxy. It therefore reports the
// modem's roaming exit rather than the host or browser's default connection.
func LookupPublicIP(ctx context.Context, networkInterface string) (PublicIPInfo, error) {
networkInterface = strings.TrimSpace(networkInterface)
if networkInterface == "" {
return PublicIPInfo{}, errors.New("cellular network interface is required")
}
if err := platformSupported(); err != nil {
return PublicIPInfo{}, err
}
dialer := boundDialer(networkInterface)
resolver := boundResolver(networkInterface)
transport := &http.Transport{
DialContext: func(ctx context.Context, _, address string) (net.Conn, error) {
return dialTarget(ctx, address, &dialer, resolver)
},
DisableKeepAlives: true,
ResponseHeaderTimeout: 12 * time.Second,
}
defer transport.CloseIdleConnections()
request, err := http.NewRequestWithContext(ctx, http.MethodGet, ipInfoURL, nil)
if err != nil {
return PublicIPInfo{}, err
}
request.Header.Set("Accept", "application/json")
request.Header.Set("User-Agent", "VoCat/1.0")
response, err := transport.RoundTrip(request)
if err != nil {
return PublicIPInfo{}, fmt.Errorf("query ipinfo.io through %s: %w", networkInterface, err)
}
defer response.Body.Close()
if response.StatusCode < 200 || response.StatusCode >= 300 {
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 4<<10))
return PublicIPInfo{}, fmt.Errorf("ipinfo.io returned HTTP %d", response.StatusCode)
}
return decodePublicIPInfo(io.LimitReader(response.Body, 64<<10))
}
func decodePublicIPInfo(reader io.Reader) (PublicIPInfo, error) {
var response struct {
IP string `json:"ip"`
Country string `json:"country"`
Region string `json:"region"`
City string `json:"city"`
Org string `json:"org"`
}
if err := json.NewDecoder(reader).Decode(&response); err != nil {
return PublicIPInfo{}, fmt.Errorf("decode ipinfo.io response: %w", err)
}
response.IP = strings.TrimSpace(response.IP)
response.Country = strings.ToUpper(strings.TrimSpace(response.Country))
if net.ParseIP(response.IP) == nil {
return PublicIPInfo{}, errors.New("ipinfo.io response contained no valid IP address")
}
if len(response.Country) != 2 {
return PublicIPInfo{}, errors.New("ipinfo.io response contained no valid country code")
}
return PublicIPInfo{
IP: response.IP, CountryCode: response.Country,
Region: strings.TrimSpace(response.Region), City: strings.TrimSpace(response.City),
Organization: strings.TrimSpace(response.Org),
}, nil
}
+22
View File
@@ -0,0 +1,22 @@
package exportproxy
import (
"strings"
"testing"
)
func TestDecodePublicIPInfo(t *testing.T) {
info, err := decodePublicIPInfo(strings.NewReader(`{"ip":"203.0.113.8","city":"London","region":"England","country":"gb","org":"AS64500 Test"}`))
if err != nil {
t.Fatal(err)
}
if info.IP != "203.0.113.8" || info.CountryCode != "GB" || info.Region != "England" || info.City != "London" {
t.Fatalf("info = %+v", info)
}
}
func TestDecodePublicIPInfoRejectsInvalidResponse(t *testing.T) {
if _, err := decodePublicIPInfo(strings.NewReader(`{"ip":"not-an-ip","country":"GB"}`)); err == nil {
t.Fatal("invalid IP was accepted")
}
}
+494
View File
@@ -0,0 +1,494 @@
package exportproxy
import (
"context"
"crypto/rand"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"log/slog"
"net"
"os"
"strconv"
"strings"
"sync"
"time"
"vocat/internal/store"
)
const (
SettingKey = "developer.export_proxy.configs"
PasswordMask = "••••••••"
ReservedID = "export-proxy"
)
var (
ErrNotFound = errors.New("export proxy configuration not found")
ErrDisabled = errors.New("export proxy is disabled")
)
type Config struct {
ID string `json:"id"`
Name string `json:"name"`
DeviceID string `json:"device_id"`
Interface string `json:"interface"`
Mode string `json:"mode"`
ListenHost string `json:"listen_host"`
ListenPort int `json:"listen_port"`
Enabled bool `json:"enabled"`
AuthEnabled bool `json:"auth_enabled"`
Username string `json:"username"`
Password string `json:"password"`
}
type Status struct {
ID string `json:"id"`
Name string `json:"name"`
Mode string `json:"mode"`
Enabled bool `json:"enabled"`
Running bool `json:"running"`
Listen string `json:"listen"`
Error string `json:"error,omitempty"`
StartedAt time.Time `json:"started_at,omitempty"`
}
type Manager struct {
mu sync.Mutex
store *store.Store
logger *slog.Logger
configs []Config
listeners map[string]net.Listener
started map[string]time.Time
lastError map[string]string
disabled bool
}
func New(ctx context.Context, database *store.Store, logger *slog.Logger, legacyConfigPath string) (*Manager, error) {
if database == nil {
return nil, errors.New("export proxy store is required")
}
if logger == nil {
logger = slog.Default()
}
manager := &Manager{
store: database, logger: logger,
listeners: make(map[string]net.Listener),
started: make(map[string]time.Time),
lastError: make(map[string]string),
}
migrated, err := manager.load(ctx, legacyConfigPath)
if err != nil {
return nil, err
}
if migrated {
if err := manager.saveLocked(ctx); err != nil {
return nil, fmt.Errorf("migrate legacy export proxy configurations: %w", err)
}
_ = RemoveLegacyConfig(legacyConfigPath)
}
for _, config := range manager.configs {
if config.Enabled {
if err := manager.start(ctx, config.ID); err != nil {
manager.logger.Warn("start built-in export proxy", "id", config.ID, "error", err)
}
}
}
return manager, nil
}
func RemoveLegacyConfig(path string) error {
path = strings.TrimSpace(path)
if path == "" {
return nil
}
err := os.Remove(path)
if errors.Is(err, os.ErrNotExist) {
return nil
}
return err
}
func (manager *Manager) load(ctx context.Context, legacyConfigPath string) (bool, error) {
setting, err := manager.store.AppSetting(ctx, SettingKey)
if err == nil {
if err := json.Unmarshal(setting.Value, &manager.configs); err != nil {
return false, fmt.Errorf("decode export proxy configurations: %w", err)
}
return false, nil
}
if !errors.Is(err, store.ErrNotFound) {
return false, err
}
legacy, err := os.ReadFile(strings.TrimSpace(legacyConfigPath))
if err != nil {
if errors.Is(err, os.ErrNotExist) || strings.TrimSpace(legacyConfigPath) == "" {
return false, nil
}
return false, err
}
if err := json.Unmarshal(legacy, &manager.configs); err != nil {
return false, fmt.Errorf("decode legacy export proxy configurations: %w", err)
}
return true, nil
}
func (manager *Manager) saveLocked(ctx context.Context) error {
raw, err := json.Marshal(manager.configs)
if err != nil {
return err
}
return manager.store.UpsertAppSetting(ctx, store.AppSetting{Key: SettingKey, Value: raw, Sensitive: true})
}
func (manager *Manager) Configs() ([]Config, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.disabled {
return nil, ErrDisabled
}
result := make([]Config, len(manager.configs))
for index, config := range manager.configs {
result[index] = redact(config)
}
return result, nil
}
// EnabledConfigForDevice returns the first enabled configuration bound to the
// given device, reporting whether one exists. It is used to block turning off a
// device's roaming data while one of its export proxies is still running.
func (manager *Manager) EnabledConfigForDevice(deviceID string) (Config, bool) {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.disabled {
return Config{}, false
}
for _, config := range manager.configs {
if config.DeviceID == deviceID && config.Enabled {
return redact(config), true
}
}
return Config{}, false
}
func (manager *Manager) Status() ([]Status, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.disabled {
return nil, ErrDisabled
}
result := make([]Status, 0, len(manager.configs))
for _, config := range manager.configs {
status := Status{ID: config.ID, Name: config.Name, Mode: config.Mode, Enabled: config.Enabled, Error: manager.lastError[config.ID]}
if listener := manager.listeners[config.ID]; listener != nil {
status.Running = true
status.Listen = listener.Addr().String()
status.StartedAt = manager.started[config.ID]
}
result = append(result, status)
}
return result, nil
}
func (manager *Manager) Create(ctx context.Context, config Config) (Config, error) {
config.ID = generateID()
if err := manager.prepareConfig(ctx, &config); err != nil {
return Config{}, err
}
manager.mu.Lock()
if manager.disabled {
manager.mu.Unlock()
return Config{}, ErrDisabled
}
if err := manager.checkPortLocked(config, ""); err != nil {
manager.mu.Unlock()
return Config{}, err
}
manager.configs = append(manager.configs, config)
if err := manager.saveLocked(ctx); err != nil {
manager.configs = manager.configs[:len(manager.configs)-1]
manager.mu.Unlock()
return Config{}, err
}
manager.mu.Unlock()
if config.Enabled {
if err := manager.start(ctx, config.ID); err != nil {
_ = manager.Delete(context.Background(), config.ID)
return Config{}, err
}
}
return redact(config), nil
}
func (manager *Manager) Update(ctx context.Context, id string, incoming Config) (Config, error) {
incoming.ID = strings.TrimSpace(id)
manager.mu.Lock()
if manager.disabled {
manager.mu.Unlock()
return Config{}, ErrDisabled
}
existing, index := manager.configByIDLocked(incoming.ID)
manager.mu.Unlock()
if index < 0 {
return Config{}, ErrNotFound
}
if incoming.Password == "" || incoming.Password == PasswordMask {
incoming.Password = existing.Password
}
if err := manager.prepareConfig(ctx, &incoming); err != nil {
return Config{}, err
}
manager.mu.Lock()
if manager.disabled {
manager.mu.Unlock()
return Config{}, ErrDisabled
}
existing, index = manager.configByIDLocked(incoming.ID)
if index < 0 {
manager.mu.Unlock()
return Config{}, ErrNotFound
}
if err := manager.checkPortLocked(incoming, incoming.ID); err != nil {
manager.mu.Unlock()
return Config{}, err
}
wasRunning := manager.listeners[incoming.ID] != nil
runtimeChanged := existing.Mode != incoming.Mode || existing.Interface != incoming.Interface ||
existing.ListenHost != incoming.ListenHost || existing.ListenPort != incoming.ListenPort ||
existing.AuthEnabled != incoming.AuthEnabled || existing.Username != incoming.Username || existing.Password != incoming.Password
manager.configs[index] = incoming
if err := manager.saveLocked(ctx); err != nil {
manager.configs[index] = existing
manager.mu.Unlock()
return Config{}, err
}
manager.mu.Unlock()
switch {
case !incoming.Enabled:
manager.stop(incoming.ID)
case !wasRunning || runtimeChanged || !existing.Enabled:
if err := manager.start(ctx, incoming.ID); err != nil {
return redact(incoming), err
}
}
return redact(incoming), nil
}
func (manager *Manager) Delete(ctx context.Context, id string) error {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.disabled {
return ErrDisabled
}
_, index := manager.configByIDLocked(strings.TrimSpace(id))
if index < 0 {
return ErrNotFound
}
manager.stopLocked(id)
manager.configs = append(manager.configs[:index], manager.configs[index+1:]...)
return manager.saveLocked(ctx)
}
// DeleteAllAndDisable is irreversible for the active developer-mode session:
// it closes every listener, removes every saved proxy, and rejects new work.
func (manager *Manager) DeleteAllAndDisable(ctx context.Context) error {
manager.mu.Lock()
for id := range manager.listeners {
manager.stopLocked(id)
}
manager.configs = nil
manager.disabled = true
manager.mu.Unlock()
err := manager.store.DeleteAppSetting(ctx, SettingKey)
if errors.Is(err, store.ErrNotFound) {
return nil
}
return err
}
func (manager *Manager) Close() error {
manager.mu.Lock()
defer manager.mu.Unlock()
manager.disabled = true
for id := range manager.listeners {
manager.stopLocked(id)
}
return nil
}
func (manager *Manager) prepareConfig(ctx context.Context, config *Config) error {
config.Name = strings.TrimSpace(config.Name)
config.DeviceID = strings.TrimSpace(config.DeviceID)
config.Interface = strings.TrimSpace(config.Interface)
config.Mode = strings.ToLower(strings.TrimSpace(config.Mode))
config.ListenHost = strings.TrimSpace(config.ListenHost)
config.Username = strings.TrimSpace(config.Username)
if config.Name == "" {
config.Name = "proxy-" + config.ID[:4]
}
if config.DeviceID == "" {
return errors.New("device is required")
}
device, err := manager.store.Device(ctx, config.DeviceID)
if err != nil {
if errors.Is(err, store.ErrNotFound) {
return errors.New("configured device was not found")
}
return err
}
if strings.TrimSpace(device.Interface) == "" {
return errors.New("the selected device has no cellular interface")
}
if config.Interface != "" && config.Interface != device.Interface {
return errors.New("proxy interface does not match the selected device")
}
config.Interface = device.Interface
if config.Enabled && !device.NetworkEnabled {
return errors.New("enable roaming data on the selected device before starting its export proxy")
}
if config.Mode != "http" && config.Mode != "socks5" {
return errors.New("mode must be http or socks5")
}
if config.ListenHost == "" {
config.ListenHost = "0.0.0.0"
}
if net.ParseIP(config.ListenHost) == nil && config.ListenHost != "localhost" {
return errors.New("listen host must be an IP address")
}
if config.ListenPort < 0 || config.ListenPort > 65535 {
return errors.New("listen port must be between 0 and 65535")
}
if config.AuthEnabled {
if config.Username == "" {
return errors.New("username is required when authentication is enabled")
}
if len(config.Username) > 128 || len(config.Password) > 128 {
return errors.New("proxy credentials are too long")
}
}
return nil
}
func (manager *Manager) checkPortLocked(config Config, excludeID string) error {
if config.ListenPort == 0 {
return nil
}
for _, current := range manager.configs {
if current.ID != excludeID && current.ListenPort == config.ListenPort && current.ListenHost == config.ListenHost {
return fmt.Errorf("port %d is already used by another export proxy", config.ListenPort)
}
}
if existing, _ := manager.configByIDLocked(excludeID); excludeID != "" &&
existing.ListenHost == config.ListenHost && existing.ListenPort == config.ListenPort {
return nil
}
listener, err := net.Listen("tcp", net.JoinHostPort(config.ListenHost, strconv.Itoa(config.ListenPort)))
if err != nil {
return fmt.Errorf("port %d is already in use", config.ListenPort)
}
_ = listener.Close()
return nil
}
func (manager *Manager) start(ctx context.Context, id string) error {
manager.mu.Lock()
if manager.disabled {
manager.mu.Unlock()
return ErrDisabled
}
config, index := manager.configByIDLocked(id)
if index < 0 || !config.Enabled {
manager.mu.Unlock()
return ErrNotFound
}
if err := platformSupported(); err != nil {
manager.lastError[id] = err.Error()
manager.mu.Unlock()
return err
}
manager.stopLocked(id)
listener, err := net.Listen("tcp", net.JoinHostPort(config.ListenHost, strconv.Itoa(config.ListenPort)))
if err != nil {
manager.lastError[id] = err.Error()
manager.mu.Unlock()
return err
}
if config.ListenPort == 0 {
config.ListenPort = listener.Addr().(*net.TCPAddr).Port
manager.configs[index] = config
if err := manager.saveLocked(ctx); err != nil {
_ = listener.Close()
manager.mu.Unlock()
return err
}
}
delete(manager.lastError, id)
manager.listeners[id] = listener
manager.started[id] = time.Now().UTC()
manager.mu.Unlock()
go manager.serve(listener, config)
return nil
}
func (manager *Manager) stop(id string) {
manager.mu.Lock()
defer manager.mu.Unlock()
manager.stopLocked(id)
}
func (manager *Manager) stopLocked(id string) {
if listener := manager.listeners[id]; listener != nil {
_ = listener.Close()
delete(manager.listeners, id)
}
delete(manager.started, id)
}
func (manager *Manager) serve(listener net.Listener, config Config) {
dialer := boundDialer(config.Interface)
resolver := boundResolver(config.Interface)
for {
connection, err := listener.Accept()
if err != nil {
return
}
go func(client net.Conn) {
defer client.Close()
var err error
if config.Mode == "http" {
err = serveHTTP(client, config, &dialer, resolver)
} else {
err = serveSOCKS(client, config, &dialer, resolver)
}
if err != nil {
manager.logger.Debug("export proxy connection closed", "id", config.ID, "error", err)
}
}(connection)
}
}
func (manager *Manager) configByIDLocked(id string) (Config, int) {
for index, config := range manager.configs {
if config.ID == id {
return config, index
}
}
return Config{}, -1
}
func redact(config Config) Config {
if config.Password != "" {
config.Password = PasswordMask
}
return config
}
func generateID() string {
value := make([]byte, 4)
_, _ = rand.Read(value)
return hex.EncodeToString(value)
}
+123
View File
@@ -0,0 +1,123 @@
package exportproxy
import (
"context"
"errors"
"io"
"log/slog"
"path/filepath"
"testing"
"vocat/internal/store"
)
func TestManagerPersistsAndDeletesDisabledConfig(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, filepath.Join(t.TempDir(), "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := database.UpsertDevice(ctx, store.Device{ID: "modem-1", Name: "modem-1", Interface: "wwan0"}); err != nil {
t.Fatal(err)
}
manager, err := New(ctx, database, slog.New(slog.NewTextHandler(io.Discard, nil)), "")
if err != nil {
t.Fatal(err)
}
created, err := manager.Create(ctx, Config{DeviceID: "modem-1", Mode: "socks5", ListenHost: "127.0.0.1", ListenPort: 1080})
if err != nil {
t.Fatal(err)
}
if created.ID == "" || created.Interface != "wwan0" {
t.Fatalf("created = %+v", created)
}
configs, err := manager.Configs()
if err != nil || len(configs) != 1 {
t.Fatalf("configs = %+v, %v", configs, err)
}
if err := manager.DeleteAllAndDisable(ctx); err != nil {
t.Fatal(err)
}
if _, err := manager.Configs(); !errors.Is(err, ErrDisabled) {
t.Fatalf("Configs after disable = %v", err)
}
if _, err := database.AppSetting(ctx, SettingKey); !errors.Is(err, store.ErrNotFound) {
t.Fatalf("setting remains: %v", err)
}
}
func TestManagerRequiresRoamingDataForEnabledProxy(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, filepath.Join(t.TempDir(), "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := database.UpsertDevice(ctx, store.Device{ID: "modem-1", Name: "modem-1", Interface: "wwan0"}); err != nil {
t.Fatal(err)
}
manager, err := New(ctx, database, nil, "")
if err != nil {
t.Fatal(err)
}
defer manager.Close()
_, err = manager.Create(ctx, Config{DeviceID: "modem-1", Mode: "socks5", ListenHost: "127.0.0.1", ListenPort: 1080, Enabled: true})
if err == nil {
t.Fatal("enabled proxy was accepted while roaming data was disabled")
}
}
func TestManagerEnabledConfigForDevice(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, filepath.Join(t.TempDir(), "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
if err := database.UpsertDevice(ctx, store.Device{ID: "modem-1", Name: "modem-1", Interface: "wwan0", NetworkEnabled: true}); err != nil {
t.Fatal(err)
}
if err := database.UpsertDevice(ctx, store.Device{ID: "modem-2", Name: "modem-2", Interface: "wwan1", NetworkEnabled: true}); err != nil {
t.Fatal(err)
}
manager, err := New(ctx, database, nil, "")
if err != nil {
t.Fatal(err)
}
defer manager.Close()
if _, ok := manager.EnabledConfigForDevice("modem-1"); ok {
t.Fatal("reported an enabled config before any was created")
}
// A disabled config bound to modem-1 must not count.
if _, err := manager.Create(ctx, Config{DeviceID: "modem-1", Mode: "socks5", ListenHost: "127.0.0.1", ListenPort: 1080}); err != nil {
t.Fatal(err)
}
if _, ok := manager.EnabledConfigForDevice("modem-1"); ok {
t.Fatal("disabled config counted as enabled")
}
// An enabled config bound to modem-2 counts only for modem-2. The listener start
// is Linux-only, so the config is created disabled and flipped on in memory to
// exercise the query without binding a port.
created, err := manager.Create(ctx, Config{DeviceID: "modem-2", Mode: "socks5", ListenHost: "127.0.0.1", ListenPort: 0, AuthEnabled: true, Username: "u", Password: "secret"})
if err != nil {
t.Fatal(err)
}
manager.mu.Lock()
for index := range manager.configs {
if manager.configs[index].ID == created.ID {
manager.configs[index].Enabled = true
}
}
manager.mu.Unlock()
if _, ok := manager.EnabledConfigForDevice("modem-1"); ok {
t.Fatal("config bound to another device counted")
}
found, ok := manager.EnabledConfigForDevice("modem-2")
if !ok {
t.Fatal("enabled config not found for its device")
}
if found.Password != PasswordMask {
t.Fatalf("password not redacted: %+v", found)
}
}
+72
View File
@@ -0,0 +1,72 @@
package exportproxy
import (
"bufio"
"context"
"encoding/base64"
"errors"
"fmt"
"net"
"net/http"
"strings"
)
func serveHTTP(client net.Conn, config Config, dialer *net.Dialer, resolver *net.Resolver) error {
reader := bufio.NewReader(client)
request, err := http.ReadRequest(reader)
if err != nil {
return err
}
if config.AuthEnabled && !httpAuthorized(request, config) {
_, _ = client.Write([]byte("HTTP/1.1 407 Proxy Authentication Required\r\nProxy-Authenticate: Basic realm=\"vocat-export-proxy\"\r\n\r\n"))
return errors.New("HTTP proxy authentication required")
}
if request.Method == http.MethodConnect {
ctx, cancel := context.WithTimeout(context.Background(), proxyTimeout)
target, err := dialTarget(ctx, request.URL.Host, dialer, resolver)
cancel()
if err != nil {
_, _ = fmt.Fprint(client, "HTTP/1.1 502 Bad Gateway\r\n\r\n")
return err
}
defer target.Close()
if _, err := client.Write([]byte("HTTP/1.1 200 Connection Established\r\n\r\n")); err != nil {
return err
}
if buffered := reader.Buffered(); buffered > 0 {
if value, err := reader.Peek(buffered); err == nil {
_, _ = target.Write(value)
_, _ = reader.Discard(buffered)
}
}
pipe(client, target)
return nil
}
request.Header.Del("Proxy-Authorization")
request.Header.Del("Proxy-Connection")
request.RequestURI = ""
transport := &http.Transport{
DialContext: func(ctx context.Context, _, address string) (net.Conn, error) {
return dialTarget(ctx, address, dialer, resolver)
},
DisableKeepAlives: true,
}
response, err := transport.RoundTrip(request)
if err != nil {
_, _ = fmt.Fprint(client, "HTTP/1.1 502 Bad Gateway\r\n\r\n")
return err
}
defer response.Body.Close()
return response.Write(client)
}
func httpAuthorized(request *http.Request, config Config) bool {
header := strings.TrimSpace(strings.TrimPrefix(request.Header.Get("Proxy-Authorization"), "Basic "))
decoded, err := base64.StdEncoding.DecodeString(header)
if err != nil {
return false
}
parts := strings.SplitN(string(decoded), ":", 2)
return len(parts) == 2 && parts[0] == config.Username && parts[1] == config.Password
}
+55
View File
@@ -0,0 +1,55 @@
package exportproxy
import (
"context"
"errors"
"fmt"
"io"
"net"
"time"
)
const proxyTimeout = 30 * time.Second
func dialTarget(ctx context.Context, address string, dialer *net.Dialer, resolver *net.Resolver) (net.Conn, error) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, err
}
if ip := net.ParseIP(host); ip != nil {
return dialer.DialContext(ctx, "tcp", net.JoinHostPort(ip.String(), port))
}
ips, err := resolver.LookupIPAddr(ctx, host)
if err != nil {
return nil, err
}
var lastErr error
for _, ip := range ips {
connection, err := dialer.DialContext(ctx, "tcp", net.JoinHostPort(ip.IP.String(), port))
if err == nil {
return connection, nil
}
lastErr = err
}
if lastErr == nil {
lastErr = fmt.Errorf("%w: no addresses for %s", errors.ErrUnsupported, host)
}
return nil, lastErr
}
func pipe(left, right net.Conn) {
done := make(chan struct{}, 2)
go func() { _, _ = copyConnection(right, left); done <- struct{}{} }()
go func() { _, _ = copyConnection(left, right); done <- struct{}{} }()
<-done
}
func copyConnection(destination net.Conn, source net.Conn) (int64, error) {
written, err := io.CopyBuffer(destination, source, make([]byte, 32*1024))
if err == nil && written > 0 {
if connection, ok := destination.(interface{ CloseWrite() error }); ok {
_ = connection.CloseWrite()
}
}
return written, err
}
+148
View File
@@ -0,0 +1,148 @@
package exportproxy
import (
"bufio"
"context"
"encoding/binary"
"errors"
"fmt"
"io"
"net"
"strconv"
)
func serveSOCKS(client net.Conn, config Config, dialer *net.Dialer, resolver *net.Resolver) error {
reader := bufio.NewReader(client)
version, err := reader.ReadByte()
if err != nil || version != 5 {
return errors.New("unsupported SOCKS version")
}
methodCount, err := reader.ReadByte()
if err != nil {
return err
}
methods := make([]byte, methodCount)
if _, err := io.ReadFull(reader, methods); err != nil {
return err
}
chosen := byte(0xff)
if config.AuthEnabled && hasMethod(methods, 2) {
chosen = 2
} else if !config.AuthEnabled && hasMethod(methods, 0) {
chosen = 0
}
if _, err := client.Write([]byte{5, chosen}); err != nil || chosen == 0xff {
return errors.New("no acceptable SOCKS authentication method")
}
if chosen == 2 {
if err := socksAuthenticate(reader, client, config); err != nil {
return err
}
}
header := make([]byte, 4)
if _, err := io.ReadFull(reader, header); err != nil {
return err
}
if header[0] != 5 || header[1] != 1 {
_ = writeSocksReply(client, 7)
return errors.New("only SOCKS5 CONNECT is supported")
}
host, port, err := readSocksAddress(reader, header[3])
if err != nil {
_ = writeSocksReply(client, 1)
return err
}
ctx, cancel := context.WithTimeout(context.Background(), proxyTimeout)
target, err := dialTarget(ctx, net.JoinHostPort(host, strconv.Itoa(port)), dialer, resolver)
cancel()
if err != nil {
_ = writeSocksReply(client, 5)
return err
}
defer target.Close()
if err := writeSocksReply(client, 0); err != nil {
return err
}
if buffered := reader.Buffered(); buffered > 0 {
if value, err := reader.Peek(buffered); err == nil {
_, _ = target.Write(value)
_, _ = reader.Discard(buffered)
}
}
pipe(client, target)
return nil
}
func socksAuthenticate(reader *bufio.Reader, writer io.Writer, config Config) error {
header := make([]byte, 2)
if _, err := io.ReadFull(reader, header); err != nil || header[0] != 1 {
return errors.New("invalid SOCKS authentication request")
}
username := make([]byte, int(header[1]))
if _, err := io.ReadFull(reader, username); err != nil {
return err
}
length, err := reader.ReadByte()
if err != nil {
return err
}
password := make([]byte, int(length))
if _, err := io.ReadFull(reader, password); err != nil {
return err
}
if string(username) != config.Username || string(password) != config.Password {
_, _ = writer.Write([]byte{1, 1})
return errors.New("SOCKS authentication failed")
}
_, err = writer.Write([]byte{1, 0})
return err
}
func hasMethod(methods []byte, wanted byte) bool {
for _, method := range methods {
if method == wanted {
return true
}
}
return false
}
func readSocksAddress(reader *bufio.Reader, kind byte) (string, int, error) {
var host string
switch kind {
case 1:
value := make([]byte, 4)
if _, err := io.ReadFull(reader, value); err != nil {
return "", 0, err
}
host = net.IP(value).String()
case 3:
length, err := reader.ReadByte()
if err != nil {
return "", 0, err
}
value := make([]byte, int(length))
if _, err := io.ReadFull(reader, value); err != nil {
return "", 0, err
}
host = string(value)
case 4:
value := make([]byte, 16)
if _, err := io.ReadFull(reader, value); err != nil {
return "", 0, err
}
host = net.IP(value).String()
default:
return "", 0, fmt.Errorf("unsupported SOCKS address type %d", kind)
}
value := make([]byte, 2)
if _, err := io.ReadFull(reader, value); err != nil {
return "", 0, err
}
return host, int(binary.BigEndian.Uint16(value)), nil
}
func writeSocksReply(connection net.Conn, code byte) error {
_, err := connection.Write([]byte{5, code, 0, 1, 0, 0, 0, 0, 0, 0})
return err
}
+527
View File
@@ -0,0 +1,527 @@
package extensions
import (
"archive/zip"
"bufio"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"mime"
"net"
"net/http"
"net/http/httputil"
"net/url"
"os"
"os/exec"
"path/filepath"
"runtime"
"sort"
"strings"
"sync"
"time"
"vocat/internal/exportproxy"
)
const maxPackageBytes int64 = 64 << 20
type Plugin struct {
Manifest
Enabled bool `json:"enabled"`
BackendAvailable bool `json:"backend_available"`
BackendRunning bool `json:"backend_running"`
BackendError string `json:"backend_error,omitempty"`
InstalledAt string `json:"installed_at"`
SHA256 string `json:"sha256"`
dir string
command *exec.Cmd
backend *url.URL
installed time.Time
}
type stateFile struct {
Enabled bool `json:"enabled"`
InstalledAt time.Time `json:"installed_at"`
SHA256 string `json:"sha256"`
}
type Manager struct {
root string
logger *slog.Logger
client *http.Client
mu sync.RWMutex
plugins map[string]*Plugin
}
func NewManager(root string, logger *slog.Logger) (*Manager, error) {
root, err := filepath.Abs(root)
if err != nil {
return nil, fmt.Errorf("resolve plugin directory: %w", err)
}
if err := os.MkdirAll(root, 0o750); err != nil {
return nil, fmt.Errorf("create plugin directory: %w", err)
}
if logger == nil {
logger = slog.Default()
}
manager := &Manager{
root: root, logger: logger, plugins: make(map[string]*Plugin),
client: &http.Client{Timeout: 45 * time.Second},
}
if err := manager.scan(); err != nil {
return nil, err
}
return manager, nil
}
func (manager *Manager) scan() error {
entries, err := os.ReadDir(manager.root)
if err != nil {
return err
}
for _, entry := range entries {
if !entry.IsDir() || !pluginIDPattern.MatchString(entry.Name()) {
continue
}
dir := filepath.Join(manager.root, entry.Name())
plugin, err := loadPlugin(dir)
if err != nil {
manager.logger.Warn("skip invalid plugin", "directory", dir, "error", err)
continue
}
if plugin.ID == exportproxy.ReservedID {
manager.logger.Info("skip legacy Export Proxy plugin; functionality is built in", "directory", dir)
continue
}
manager.plugins[plugin.ID] = plugin
if plugin.Enabled {
manager.startLocked(plugin)
}
}
return nil
}
func loadPlugin(dir string) (*Plugin, error) {
file, err := os.Open(filepath.Join(dir, ManifestFilename))
if err != nil {
return nil, err
}
manifest, err := DecodeManifest(file)
_ = file.Close()
if err != nil {
return nil, err
}
if filepath.Base(dir) != manifest.ID {
return nil, errors.New("plugin directory does not match manifest id")
}
var state stateFile
stateData, err := os.ReadFile(filepath.Join(dir, ".vocat-state.json"))
if err == nil {
if err := json.Unmarshal(stateData, &state); err != nil {
return nil, fmt.Errorf("decode plugin state: %w", err)
}
}
command, available := manifest.BackendCommand()
if available {
_, err = os.Stat(filepath.Join(dir, filepath.FromSlash(command)))
available = err == nil
}
return &Plugin{
Manifest: manifest, Enabled: state.Enabled, BackendAvailable: available,
InstalledAt: state.InstalledAt.UTC().Format(time.RFC3339), SHA256: state.SHA256,
dir: dir, installed: state.InstalledAt,
}, nil
}
func (manager *Manager) List() []Plugin {
manager.mu.RLock()
defer manager.mu.RUnlock()
result := make([]Plugin, 0, len(manager.plugins))
for _, plugin := range manager.plugins {
copy := *plugin
copy.command = nil
copy.backend = nil
result = append(result, copy)
}
sort.Slice(result, func(i, j int) bool { return result[i].Name < result[j].Name })
return result
}
func (manager *Manager) InstallURL(ctx context.Context, rawURL, expectedSHA string) (Plugin, error) {
parsed, err := url.Parse(strings.TrimSpace(rawURL))
if err != nil || (parsed.Scheme != "https" && parsed.Scheme != "http") || parsed.Host == "" {
return Plugin{}, errors.New("plugin URL must be an absolute HTTP or HTTPS URL")
}
request, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
if err != nil {
return Plugin{}, err
}
response, err := manager.client.Do(request)
if err != nil {
return Plugin{}, fmt.Errorf("download plugin: %w", err)
}
defer response.Body.Close()
if response.StatusCode < 200 || response.StatusCode >= 300 {
return Plugin{}, fmt.Errorf("download plugin: HTTP %d", response.StatusCode)
}
return manager.Install(response.Body, expectedSHA)
}
func (manager *Manager) Install(reader io.Reader, expectedSHA string) (Plugin, error) {
temp, err := os.CreateTemp(manager.root, ".upload-*.vocat-plugin")
if err != nil {
return Plugin{}, err
}
tempName := temp.Name()
defer os.Remove(tempName)
hash := sha256.New()
written, copyErr := io.Copy(io.MultiWriter(temp, hash), io.LimitReader(reader, maxPackageBytes+1))
closeErr := temp.Close()
if copyErr != nil {
return Plugin{}, copyErr
}
if closeErr != nil {
return Plugin{}, closeErr
}
if written > maxPackageBytes {
return Plugin{}, fmt.Errorf("plugin package exceeds %d MiB", maxPackageBytes>>20)
}
actualSHA := hex.EncodeToString(hash.Sum(nil))
if expected := strings.ToLower(strings.TrimSpace(expectedSHA)); expected != "" && expected != actualSHA {
return Plugin{}, errors.New("plugin package SHA-256 does not match")
}
archive, err := zip.OpenReader(tempName)
if err != nil {
return Plugin{}, errors.New("plugin package must be a ZIP archive")
}
defer archive.Close()
manifest, err := manifestFromArchive(archive.File)
if err != nil {
return Plugin{}, err
}
if manifest.ID == exportproxy.ReservedID {
return Plugin{}, errors.New("plugin ID export-proxy is reserved by the built-in Export Proxy feature")
}
staging, err := os.MkdirTemp(manager.root, ".install-"+manifest.ID+"-")
if err != nil {
return Plugin{}, err
}
defer os.RemoveAll(staging)
if err := extractArchive(archive.File, staging); err != nil {
return Plugin{}, err
}
installedAt := time.Now().UTC()
state := stateFile{Enabled: true, InstalledAt: installedAt, SHA256: actualSHA}
if err := writeState(staging, state); err != nil {
return Plugin{}, err
}
target := filepath.Join(manager.root, manifest.ID)
manager.mu.Lock()
defer manager.mu.Unlock()
if _, exists := manager.plugins[manifest.ID]; exists {
return Plugin{}, fmt.Errorf("plugin %q is already installed; uninstall it before replacing", manifest.ID)
}
if err := os.Rename(staging, target); err != nil {
return Plugin{}, fmt.Errorf("activate plugin: %w", err)
}
plugin, err := loadPlugin(target)
if err != nil {
_ = os.RemoveAll(target)
return Plugin{}, err
}
manager.plugins[plugin.ID] = plugin
manager.startLocked(plugin)
return publicPlugin(plugin), nil
}
func manifestFromArchive(files []*zip.File) (Manifest, error) {
for _, file := range files {
name := strings.ReplaceAll(file.Name, `\`, "/")
if name != ManifestFilename {
continue
}
reader, err := file.Open()
if err != nil {
return Manifest{}, err
}
manifest, decodeErr := DecodeManifest(reader)
_ = reader.Close()
return manifest, decodeErr
}
return Manifest{}, fmt.Errorf("plugin package is missing root %s", ManifestFilename)
}
func extractArchive(files []*zip.File, staging string) error {
var expanded int64
for _, file := range files {
name := strings.ReplaceAll(file.Name, `\`, "/")
if strings.HasSuffix(name, "/") {
continue
}
if !safeRelativePath(name) || file.Mode()&os.ModeSymlink != 0 {
return fmt.Errorf("plugin package contains unsafe path %q", file.Name)
}
expanded += int64(file.UncompressedSize64)
if expanded > maxPackageBytes*4 {
return errors.New("expanded plugin package is too large")
}
target := filepath.Join(staging, filepath.FromSlash(name))
if err := os.MkdirAll(filepath.Dir(target), 0o750); err != nil {
return err
}
input, err := file.Open()
if err != nil {
return err
}
mode := os.FileMode(0o640)
if file.Mode()&0o111 != 0 {
mode = 0o750
}
output, err := os.OpenFile(target, os.O_CREATE|os.O_EXCL|os.O_WRONLY, mode)
if err == nil {
_, err = io.Copy(output, input)
_ = output.Close()
}
_ = input.Close()
if err != nil {
return err
}
}
return nil
}
func writeState(dir string, state stateFile) error {
data, err := json.MarshalIndent(state, "", " ")
if err != nil {
return err
}
return os.WriteFile(filepath.Join(dir, ".vocat-state.json"), data, 0o600)
}
func (manager *Manager) SetEnabled(id string, enabled bool) (Plugin, error) {
manager.mu.Lock()
defer manager.mu.Unlock()
plugin := manager.plugins[id]
if plugin == nil {
return Plugin{}, os.ErrNotExist
}
if plugin.Enabled == enabled {
return publicPlugin(plugin), nil
}
plugin.Enabled = enabled
if err := writeState(plugin.dir, stateFile{Enabled: enabled, InstalledAt: plugin.installed, SHA256: plugin.SHA256}); err != nil {
plugin.Enabled = !enabled
return Plugin{}, err
}
if enabled {
manager.startLocked(plugin)
} else {
manager.stopLocked(plugin)
}
return publicPlugin(plugin), nil
}
func (manager *Manager) Uninstall(id string) error {
manager.mu.Lock()
defer manager.mu.Unlock()
plugin := manager.plugins[id]
if plugin == nil {
return os.ErrNotExist
}
manager.stopLocked(plugin)
delete(manager.plugins, id)
clean := filepath.Clean(plugin.dir)
if filepath.Dir(clean) != filepath.Clean(manager.root) || filepath.Base(clean) != id {
return errors.New("refusing to remove plugin outside plugin directory")
}
return os.RemoveAll(clean)
}
func (manager *Manager) ServeAsset(w http.ResponseWriter, r *http.Request, id, name string) {
manager.mu.RLock()
plugin := manager.plugins[id]
manager.mu.RUnlock()
if plugin == nil || !plugin.Enabled || !safeRelativePath(name) {
http.NotFound(w, r)
return
}
filename := filepath.Join(plugin.dir, filepath.FromSlash(name))
if !strings.HasPrefix(filepath.Clean(filename), filepath.Clean(plugin.dir)+string(os.PathSeparator)) {
http.NotFound(w, r)
return
}
file, err := os.Open(filename)
if err != nil {
http.NotFound(w, r)
return
}
defer file.Close()
info, err := file.Stat()
if err != nil || !info.Mode().IsRegular() {
http.NotFound(w, r)
return
}
contentType := mime.TypeByExtension(filepath.Ext(filename))
if contentType != "" {
w.Header().Set("Content-Type", contentType)
}
w.Header().Set("Cache-Control", "no-cache")
http.ServeContent(w, r, info.Name(), info.ModTime(), file)
}
func (manager *Manager) ProxyBackend(w http.ResponseWriter, r *http.Request, id string) {
manager.mu.RLock()
plugin := manager.plugins[id]
var target *url.URL
if plugin != nil && plugin.Enabled && plugin.BackendRunning && plugin.backend != nil {
copy := *plugin.backend
target = &copy
}
manager.mu.RUnlock()
if target == nil {
http.Error(w, "plugin backend is not running", http.StatusServiceUnavailable)
return
}
proxy := httputil.NewSingleHostReverseProxy(target)
originalDirector := proxy.Director
proxy.Director = func(request *http.Request) {
originalDirector(request)
request.Header.Del("Cookie")
request.Header.Del("X-CSRF-Token")
request.Header.Del("Authorization")
request.Header.Set("X-VoCat-Plugin-ID", id)
}
proxy.ModifyResponse = func(response *http.Response) error {
response.Header.Del("Set-Cookie")
return nil
}
prefix := "/api/extensions/" + id + "/backend"
r.URL.Path = strings.TrimPrefix(r.URL.Path, prefix)
if r.URL.Path == "" {
r.URL.Path = "/"
}
proxy.ErrorHandler = func(w http.ResponseWriter, _ *http.Request, err error) {
manager.logger.Warn("plugin backend proxy failed", "plugin", id, "error", err)
http.Error(w, "plugin backend is unavailable", http.StatusBadGateway)
}
proxy.ServeHTTP(w, r)
}
func (manager *Manager) startLocked(plugin *Plugin) {
commandPath, supported := plugin.BackendCommand()
if !supported {
if plugin.Backend != nil {
plugin.BackendError = "backend does not support " + runtime.GOOS + "/" + runtime.GOARCH
}
return
}
fullCommand := filepath.Join(plugin.dir, filepath.FromSlash(commandPath))
if _, err := os.Stat(fullCommand); err != nil {
plugin.BackendError = "backend executable is missing"
return
}
if runtime.GOOS != "windows" {
if err := os.Chmod(fullCommand, 0o750); err != nil {
plugin.BackendError = "make backend executable: " + err.Error()
return
}
}
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
plugin.BackendError = err.Error()
return
}
address := listener.Addr().String()
_ = listener.Close()
dataDir := filepath.Join(plugin.dir, "data")
if err := os.MkdirAll(dataDir, 0o700); err != nil {
plugin.BackendError = err.Error()
return
}
command := exec.Command(fullCommand)
command.Dir = plugin.dir
command.Env = append(os.Environ(),
"VOCAT_PLUGIN_ID="+plugin.ID,
"VOCAT_PLUGIN_LISTEN="+address,
"VOCAT_PLUGIN_DATA_DIR="+dataDir,
)
stdout, err := command.StdoutPipe()
if err != nil {
plugin.BackendError = err.Error()
return
}
command.Stderr = command.Stdout
if err := command.Start(); err != nil {
plugin.BackendError = err.Error()
return
}
plugin.command = command
plugin.backend = &url.URL{Scheme: "http", Host: address}
plugin.BackendRunning = true
plugin.BackendError = ""
go manager.captureOutput(plugin.ID, stdout)
go manager.waitProcess(plugin.ID, command)
}
func (manager *Manager) captureOutput(id string, reader io.Reader) {
scanner := bufio.NewScanner(reader)
for scanner.Scan() {
manager.logger.Info("plugin output", "plugin", id, "message", scanner.Text())
}
}
func (manager *Manager) waitProcess(id string, command *exec.Cmd) {
err := command.Wait()
manager.mu.Lock()
defer manager.mu.Unlock()
plugin := manager.plugins[id]
if plugin == nil || plugin.command != command {
return
}
plugin.command = nil
plugin.backend = nil
plugin.BackendRunning = false
if err != nil && plugin.Enabled {
plugin.BackendError = err.Error()
manager.logger.Warn("plugin backend exited", "plugin", id, "error", err)
}
}
func (manager *Manager) stopLocked(plugin *Plugin) {
command := plugin.command
plugin.command = nil
plugin.backend = nil
plugin.BackendRunning = false
if command != nil && command.Process != nil {
_ = command.Process.Signal(os.Interrupt)
go func() {
timer := time.NewTimer(3 * time.Second)
defer timer.Stop()
<-timer.C
_ = command.Process.Kill()
}()
}
}
func (manager *Manager) Close() {
manager.mu.Lock()
defer manager.mu.Unlock()
for _, plugin := range manager.plugins {
manager.stopLocked(plugin)
}
}
func publicPlugin(plugin *Plugin) Plugin {
copy := *plugin
copy.command = nil
copy.backend = nil
return copy
}
+100
View File
@@ -0,0 +1,100 @@
package extensions
import (
"archive/zip"
"bytes"
"io"
"log/slog"
"strings"
"testing"
)
func TestInstallListDisableAndUninstall(t *testing.T) {
manager, err := NewManager(t.TempDir(), slog.New(slog.NewTextHandler(io.Discard, nil)))
if err != nil {
t.Fatal(err)
}
defer manager.Close()
archive := testPackage(t, map[string]string{
ManifestFilename: `{
"schema_version":1,"id":"test-plugin","name":"Test","version":"1.0.0",
"permissions":["devices.read"],
"contributions":[{"id":"test-page","label":"Test","location":"sidebar","entry":"web/index.html"}]
}`,
"web/index.html": "<h1>test</h1>",
})
plugin, err := manager.Install(bytes.NewReader(archive), "")
if err != nil {
t.Fatal(err)
}
if plugin.ID != "test-plugin" || !plugin.Enabled || len(manager.List()) != 1 {
t.Fatalf("unexpected installed plugin: %#v", plugin)
}
if _, err := manager.SetEnabled(plugin.ID, false); err != nil {
t.Fatal(err)
}
if manager.List()[0].Enabled {
t.Fatal("plugin remained enabled")
}
if err := manager.Uninstall(plugin.ID); err != nil {
t.Fatal(err)
}
if len(manager.List()) != 0 {
t.Fatal("plugin remained installed")
}
}
func TestInstallRejectsPathTraversal(t *testing.T) {
manager, err := NewManager(t.TempDir(), nil)
if err != nil {
t.Fatal(err)
}
defer manager.Close()
archive := testPackage(t, map[string]string{
ManifestFilename: `{
"schema_version":1,"id":"bad-plugin","name":"Bad","version":"1",
"contributions":[{"id":"bad-page","label":"Bad","location":"sidebar","entry":"web/index.html"}]
}`,
"../escaped": "bad",
})
if _, err := manager.Install(bytes.NewReader(archive), ""); err == nil || !strings.Contains(err.Error(), "unsafe path") {
t.Fatalf("Install traversal error = %v", err)
}
}
func TestInstallVerifiesSHA256(t *testing.T) {
manager, err := NewManager(t.TempDir(), nil)
if err != nil {
t.Fatal(err)
}
defer manager.Close()
archive := testPackage(t, map[string]string{
ManifestFilename: `{
"schema_version":1,"id":"hash-plugin","name":"Hash","version":"1",
"contributions":[{"id":"hash-page","label":"Hash","location":"sidebar","entry":"web/index.html"}]
}`,
"web/index.html": "ok",
})
if _, err := manager.Install(bytes.NewReader(archive), strings.Repeat("0", 64)); err == nil {
t.Fatal("Install accepted incorrect SHA-256")
}
}
func testPackage(t *testing.T, files map[string]string) []byte {
t.Helper()
var output bytes.Buffer
writer := zip.NewWriter(&output)
for name, content := range files {
entry, err := writer.Create(name)
if err != nil {
t.Fatal(err)
}
if _, err := entry.Write([]byte(content)); err != nil {
t.Fatal(err)
}
}
if err := writer.Close(); err != nil {
t.Fatal(err)
}
return output.Bytes()
}
+135
View File
@@ -0,0 +1,135 @@
package extensions
import (
"encoding/json"
"errors"
"fmt"
"io"
"regexp"
"runtime"
"sort"
"strings"
)
const (
ManifestFilename = "vocat-plugin.json"
SchemaVersion = 1
)
var pluginIDPattern = regexp.MustCompile(`^[a-z][a-z0-9-]{1,62}[a-z0-9]$`)
type Contribution struct {
ID string `json:"id"`
Label string `json:"label"`
LabelZH string `json:"label_zh,omitempty"`
Location string `json:"location"`
After string `json:"after,omitempty"`
Entry string `json:"entry"`
}
type Backend struct {
Commands map[string]string `json:"commands,omitempty"`
}
type Manifest struct {
SchemaVersion int `json:"schema_version"`
ID string `json:"id"`
Name string `json:"name"`
Version string `json:"version"`
Description string `json:"description,omitempty"`
Author string `json:"author,omitempty"`
Homepage string `json:"homepage,omitempty"`
Permissions []string `json:"permissions,omitempty"`
Contributions []Contribution `json:"contributions"`
Backend *Backend `json:"backend,omitempty"`
}
func DecodeManifest(reader io.Reader) (Manifest, error) {
decoder := json.NewDecoder(io.LimitReader(reader, 256<<10))
decoder.DisallowUnknownFields()
var manifest Manifest
if err := decoder.Decode(&manifest); err != nil {
return Manifest{}, fmt.Errorf("decode %s: %w", ManifestFilename, err)
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return Manifest{}, errors.New("plugin manifest must contain one JSON object")
}
if err := manifest.Validate(); err != nil {
return Manifest{}, err
}
return manifest, nil
}
func (manifest Manifest) Validate() error {
if manifest.SchemaVersion != SchemaVersion {
return fmt.Errorf("unsupported plugin schema_version %d", manifest.SchemaVersion)
}
if !pluginIDPattern.MatchString(manifest.ID) {
return errors.New("plugin id must be 3-64 lowercase letters, digits, or hyphens")
}
if strings.TrimSpace(manifest.Name) == "" || len(manifest.Name) > 100 {
return errors.New("plugin name is required and must not exceed 100 characters")
}
if strings.TrimSpace(manifest.Version) == "" || len(manifest.Version) > 64 {
return errors.New("plugin version is required and must not exceed 64 characters")
}
seen := make(map[string]struct{}, len(manifest.Contributions))
for _, contribution := range manifest.Contributions {
if !pluginIDPattern.MatchString(contribution.ID) {
return fmt.Errorf("invalid contribution id %q", contribution.ID)
}
if _, duplicate := seen[contribution.ID]; duplicate {
return fmt.Errorf("duplicate contribution id %q", contribution.ID)
}
seen[contribution.ID] = struct{}{}
if contribution.Location != "sidebar" && contribution.Location != "proxy" {
return fmt.Errorf("contribution %q has unsupported location %q", contribution.ID, contribution.Location)
}
if strings.TrimSpace(contribution.Label) == "" {
return fmt.Errorf("contribution %q requires a label", contribution.ID)
}
if !safeRelativePath(contribution.Entry) {
return fmt.Errorf("contribution %q has an unsafe entry path", contribution.ID)
}
}
if manifest.Backend != nil {
if len(manifest.Backend.Commands) == 0 {
return errors.New("plugin backend commands are empty")
}
for platform, command := range manifest.Backend.Commands {
if !strings.Contains(platform, "/") || !safeRelativePath(command) {
return fmt.Errorf("plugin backend command for %q is invalid", platform)
}
}
}
permissions := append([]string(nil), manifest.Permissions...)
sort.Strings(permissions)
for index, permission := range permissions {
if strings.TrimSpace(permission) == "" || (index > 0 && permission == permissions[index-1]) {
return errors.New("plugin permissions must be non-empty and unique")
}
}
return nil
}
func (manifest Manifest) BackendCommand() (string, bool) {
if manifest.Backend == nil {
return "", false
}
command, ok := manifest.Backend.Commands[runtime.GOOS+"/"+runtime.GOARCH]
return command, ok
}
func safeRelativePath(value string) bool {
value = strings.ReplaceAll(strings.TrimSpace(value), `\`, "/")
if value == "" || strings.HasPrefix(value, "/") || strings.Contains(value, ":") {
return false
}
for _, segment := range strings.Split(value, "/") {
if segment == "" || segment == "." || segment == ".." {
return false
}
}
return true
}
+106
View File
@@ -0,0 +1,106 @@
package httpsmode
import (
"bufio"
"errors"
"net"
"sync"
"time"
)
type bufferedConn struct {
net.Conn
reader *bufio.Reader
}
func (conn *bufferedConn) Read(buffer []byte) (int, error) { return conn.reader.Read(buffer) }
type channelListener struct {
address net.Addr
conns chan net.Conn
done chan struct{}
}
func (listener *channelListener) Accept() (net.Conn, error) {
select {
case conn := <-listener.conns:
if conn == nil {
return nil, net.ErrClosed
}
return conn, nil
case <-listener.done:
return nil, net.ErrClosed
}
}
func (listener *channelListener) Close() error { return nil }
func (listener *channelListener) Addr() net.Addr { return listener.address }
type Multiplexer struct {
base net.Listener
manager *Manager
plain *channelListener
tls *channelListener
done chan struct{}
closeOnce sync.Once
}
func NewMultiplexer(base net.Listener, manager *Manager) *Multiplexer {
done := make(chan struct{})
mux := &Multiplexer{
base: base, manager: manager, done: done,
plain: &channelListener{address: base.Addr(), conns: make(chan net.Conn, 64), done: done},
tls: &channelListener{address: base.Addr(), conns: make(chan net.Conn, 64), done: done},
}
go mux.accept()
return mux
}
func (mux *Multiplexer) Plain() net.Listener { return mux.plain }
func (mux *Multiplexer) TLS() net.Listener { return mux.tls }
func (mux *Multiplexer) Close() error {
var err error
mux.closeOnce.Do(func() {
close(mux.done)
err = mux.base.Close()
})
return err
}
func (mux *Multiplexer) accept() {
for {
conn, err := mux.base.Accept()
if err != nil {
if !errors.Is(err, net.ErrClosed) {
_ = mux.Close()
}
return
}
go mux.classify(conn)
}
}
func (mux *Multiplexer) classify(conn net.Conn) {
reader := bufio.NewReaderSize(conn, 4096)
_ = conn.SetReadDeadline(time.Now().Add(10 * time.Second))
first, err := reader.Peek(1)
_ = conn.SetReadDeadline(time.Time{})
if err != nil {
_ = conn.Close()
return
}
wrapped := &bufferedConn{Conn: conn, reader: reader}
listener := mux.plain
if first[0] == 0x16 {
if !mux.manager.Enabled() {
_ = conn.Close()
return
}
listener = mux.tls
}
select {
case listener.conns <- wrapped:
case <-mux.done:
_ = conn.Close()
}
}
+259
View File
@@ -0,0 +1,259 @@
package httpsmode
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/sha256"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/hex"
"encoding/json"
"encoding/pem"
"errors"
"fmt"
"math/big"
"net"
"os"
"path/filepath"
"strings"
"sync"
"sync/atomic"
"time"
"vocat/internal/store"
)
const SettingKey = "transport.self_signed_https"
type State struct {
Enabled bool `json:"enabled"`
HTTPURL string `json:"http_url"`
HTTPSURL string `json:"https_url"`
Fingerprint string `json:"fingerprint,omitempty"`
NotAfter time.Time `json:"not_after,omitempty"`
}
type Manager struct {
store *store.Store
dir string
address string
enabled atomic.Bool
mu sync.RWMutex
cert *tls.Certificate
}
func New(ctx context.Context, database *store.Store, dir, address string) (*Manager, error) {
manager := &Manager{store: database, dir: dir, address: address}
setting, err := database.AppSetting(ctx, SettingKey)
if err == nil {
var document struct {
Enabled bool `json:"enabled"`
}
if json.Unmarshal(setting.Value, &document) == nil && document.Enabled {
if err := manager.ensureCertificate(); err != nil {
return nil, err
}
manager.enabled.Store(true)
}
} else if !errors.Is(err, store.ErrNotFound) {
return nil, err
}
return manager, nil
}
func (manager *Manager) Enabled() bool { return manager != nil && manager.enabled.Load() }
func (manager *Manager) SetEnabled(ctx context.Context, enabled bool) (State, error) {
if enabled {
if err := manager.ensureCertificate(); err != nil {
return State{}, err
}
}
raw, err := json.Marshal(map[string]bool{"enabled": enabled})
if err != nil {
return State{}, err
}
if err := manager.store.UpsertAppSetting(ctx, store.AppSetting{Key: SettingKey, Value: raw}); err != nil {
return State{}, err
}
manager.enabled.Store(enabled)
return manager.State(""), nil
}
func (manager *Manager) State(host string) State {
host = strings.TrimSpace(host)
if host == "" {
host = manager.address
}
state := State{
Enabled: manager.Enabled(),
HTTPURL: "http://" + host,
HTTPSURL: "https://" + host,
}
manager.mu.RLock()
if manager.cert != nil && manager.cert.Leaf != nil {
digest := sha256.Sum256(manager.cert.Leaf.Raw)
encoded := strings.ToUpper(hex.EncodeToString(digest[:]))
parts := make([]string, 0, len(encoded)/2)
for len(encoded) >= 2 {
parts = append(parts, encoded[:2])
encoded = encoded[2:]
}
state.Fingerprint = strings.Join(parts, ":")
state.NotAfter = manager.cert.Leaf.NotAfter
}
manager.mu.RUnlock()
return state
}
func (manager *Manager) TLSConfig() *tls.Config {
return &tls.Config{
MinVersion: tls.VersionTLS12,
NextProtos: []string{"h2", "http/1.1"},
GetCertificate: func(*tls.ClientHelloInfo) (*tls.Certificate, error) {
manager.mu.RLock()
defer manager.mu.RUnlock()
if manager.cert == nil {
return nil, errors.New("self-signed certificate is unavailable")
}
return manager.cert, nil
},
}
}
func (manager *Manager) CertificatePEM() ([]byte, error) {
if err := manager.ensureCertificate(); err != nil {
return nil, err
}
return os.ReadFile(filepath.Join(manager.dir, "selfsigned.crt"))
}
func (manager *Manager) ensureCertificate() error {
manager.mu.Lock()
defer manager.mu.Unlock()
if manager.cert != nil && manager.cert.Leaf != nil && time.Until(manager.cert.Leaf.NotAfter) > 30*24*time.Hour {
return nil
}
if err := os.MkdirAll(manager.dir, 0o750); err != nil {
return fmt.Errorf("create TLS directory: %w", err)
}
certPath := filepath.Join(manager.dir, "selfsigned.crt")
keyPath := filepath.Join(manager.dir, "selfsigned.key")
if cert, err := loadCertificate(certPath, keyPath); err == nil && time.Until(cert.Leaf.NotAfter) > 30*24*time.Hour {
manager.cert = cert
return nil
}
certPEM, keyPEM, err := generateCertificate(manager.address)
if err != nil {
return err
}
if err := writePrivateFile(keyPath, keyPEM, 0o600); err != nil {
return err
}
if err := writePrivateFile(certPath, certPEM, 0o644); err != nil {
return err
}
cert, err := loadCertificate(certPath, keyPath)
if err != nil {
return err
}
manager.cert = cert
return nil
}
func loadCertificate(certPath, keyPath string) (*tls.Certificate, error) {
cert, err := tls.LoadX509KeyPair(certPath, keyPath)
if err != nil {
return nil, err
}
if len(cert.Certificate) == 0 {
return nil, errors.New("certificate chain is empty")
}
cert.Leaf, err = x509.ParseCertificate(cert.Certificate[0])
if err != nil {
return nil, err
}
return &cert, nil
}
func generateCertificate(address string) ([]byte, []byte, error) {
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
return nil, nil, err
}
limit := new(big.Int).Lsh(big.NewInt(1), 128)
serial, err := rand.Int(rand.Reader, limit)
if err != nil {
return nil, nil, err
}
now := time.Now().UTC()
template := &x509.Certificate{
SerialNumber: serial,
Subject: pkix.Name{CommonName: "VoCat self-signed local certificate", Organization: []string{"VoCat"}},
NotBefore: now.Add(-5 * time.Minute), NotAfter: now.AddDate(5, 0, 0),
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
DNSNames: []string{"localhost"},
IPAddresses: []net.IP{net.IPv4(127, 0, 0, 1), net.IPv6loopback},
}
if hostname, hostnameErr := os.Hostname(); hostnameErr == nil && strings.TrimSpace(hostname) != "" {
template.DNSNames = append(template.DNSNames, strings.TrimSpace(hostname))
}
if host, _, splitErr := net.SplitHostPort(address); splitErr == nil {
if ip := net.ParseIP(host); ip != nil && !ip.IsUnspecified() {
template.IPAddresses = append(template.IPAddresses, ip)
} else if host != "" && host != "0.0.0.0" && host != "::" {
template.DNSNames = append(template.DNSNames, host)
}
}
if interfaces, interfaceErr := net.InterfaceAddrs(); interfaceErr == nil {
for _, item := range interfaces {
text := item.String()
if slash := strings.IndexByte(text, '/'); slash >= 0 {
text = text[:slash]
}
if ip := net.ParseIP(strings.TrimSpace(text)); ip != nil && !ip.IsUnspecified() {
template.IPAddresses = append(template.IPAddresses, ip)
}
}
}
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
if err != nil {
return nil, nil, err
}
keyDER, err := x509.MarshalPKCS8PrivateKey(key)
if err != nil {
return nil, nil, err
}
return pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der}),
pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: keyDER}), nil
}
func writePrivateFile(path string, data []byte, mode os.FileMode) error {
temp, err := os.CreateTemp(filepath.Dir(path), ".tls-*")
if err != nil {
return err
}
tempName := temp.Name()
defer os.Remove(tempName)
if err := temp.Chmod(mode); err != nil {
_ = temp.Close()
return err
}
if _, err := temp.Write(data); err != nil {
_ = temp.Close()
return err
}
if err := temp.Sync(); err != nil {
_ = temp.Close()
return err
}
if err := temp.Close(); err != nil {
return err
}
return os.Rename(tempName, path)
}
+99
View File
@@ -0,0 +1,99 @@
package httpsmode
import (
"context"
"crypto/tls"
"net"
"path/filepath"
"testing"
"vocat/internal/store"
)
func TestManagerPersistsToggleAndCertificate(t *testing.T) {
ctx := context.Background()
dir := t.TempDir()
database, err := store.Open(ctx, filepath.Join(dir, "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
manager, err := New(ctx, database, filepath.Join(dir, "tls"), "0.0.0.0:7575")
if err != nil {
t.Fatal(err)
}
state, err := manager.SetEnabled(ctx, true)
if err != nil {
t.Fatal(err)
}
if !state.Enabled || state.Fingerprint == "" || state.NotAfter.IsZero() {
t.Fatalf("enabled state = %#v", state)
}
certificate, err := manager.CertificatePEM()
if err != nil || len(certificate) == 0 {
t.Fatalf("certificate = %d bytes, %v", len(certificate), err)
}
reloaded, err := New(ctx, database, filepath.Join(dir, "tls"), "0.0.0.0:7575")
if err != nil || !reloaded.Enabled() {
t.Fatalf("reloaded manager enabled=%v error=%v", reloaded.Enabled(), err)
}
if _, err := reloaded.SetEnabled(ctx, false); err != nil || reloaded.Enabled() {
t.Fatalf("disable enabled=%v error=%v", reloaded.Enabled(), err)
}
}
func TestMultiplexerRoutesPlainAndTLS(t *testing.T) {
ctx := context.Background()
dir := t.TempDir()
database, err := store.Open(ctx, filepath.Join(dir, "vocat.db"))
if err != nil {
t.Fatal(err)
}
defer database.Close()
manager, err := New(ctx, database, filepath.Join(dir, "tls"), "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
if _, err := manager.SetEnabled(ctx, true); err != nil {
t.Fatal(err)
}
base, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatal(err)
}
mux := NewMultiplexer(base, manager)
defer mux.Close()
plainClient, err := net.Dial("tcp", base.Addr().String())
if err != nil {
t.Fatal(err)
}
defer plainClient.Close()
if _, err := plainClient.Write([]byte("GET / HTTP/1.1\r\nHost: local\r\n\r\n")); err != nil {
t.Fatal(err)
}
plainServer, err := mux.Plain().Accept()
if err != nil {
t.Fatal(err)
}
defer plainServer.Close()
tlsResult := make(chan error, 1)
go func() {
serverConn, acceptErr := mux.TLS().Accept()
if acceptErr != nil {
tlsResult <- acceptErr
return
}
defer serverConn.Close()
tlsResult <- tls.Server(serverConn, manager.TLSConfig()).Handshake()
}()
tlsClient, err := tls.Dial("tcp", base.Addr().String(), &tls.Config{InsecureSkipVerify: true}) // test-only local certificate
if err != nil {
t.Fatal(err)
}
_ = tlsClient.Close()
if err := <-tlsResult; err != nil {
t.Fatal(err)
}
}
+1
View File
@@ -74,6 +74,7 @@ var zhToEn = map[string]string{
// ---- devices ----
"设备数量已达上限,最多只能添加 %d 台设备": "Device limit reached; at most %d devices can be added.",
"SIM 卡归属地为%s(MCC %s),本服务不向该地区卡片提供数据/短信/VoWiFi": "The SIM's home region is %s (MCC %s); this service does not provide data, SMS, or VoWiFi to cards from that region.",
"请先禁用该设备已绑定的导出代理,再关闭漫游数据": "Disable the export proxy bound to this device before turning off roaming data.",
// ---- settings / update ----
"未配置受信任的软件更新源;不会从未知地址下载或执行文件。": "No trusted update source is configured; no files will be downloaded or executed from unknown addresses.",
+37 -21
View File
@@ -127,6 +127,7 @@ func (d *SysFSDiscoverer) Discover(ctx context.Context) ([]Candidate, error) {
}
return left.Name < right.Name
})
assignQuectelPortRoles(state.candidate.Ports)
state.candidate.ATPort = selectATPort(state.candidate.Ports)
result = append(result, state.candidate)
}
@@ -240,30 +241,45 @@ func sanitizeID(value string) string {
return strings.Trim(result.String(), "-")
}
func quecPortRole(interfaceNumber int, name string) PortRole {
// Quectel exposes the same logical ports under more than one USB
// composition. In both layouts seen on EC20/EC25 hardware the kernel
// stable tty name is the stronger hint: ttyUSB0 is diagnostic and
// ttyUSB2 is the primary AT port, even when their interface numbers are
// 00/02 instead of 02/04.
switch name {
case "ttyUSB0":
return PortRoleDiagnostic
case "ttyUSB1":
return PortRoleNMEA
case "ttyUSB2":
return PortRoleAT
case "ttyUSB3":
return PortRoleModem
func assignQuectelPortRoles(ports []Port) {
// ttyUSB numbers are allocated globally by Linux. A second modem therefore
// commonly exposes ttyUSB4..ttyUSB7, so absolute tty names cannot identify
// the logical AT port. Infer the Quectel composition once per physical USB
// device and assign roles from that device's interface numbers.
base := 0x02
for _, port := range ports {
if port.InterfaceNumber <= 0x01 {
base = 0x00
break
}
}
for index := range ports {
switch ports[index].InterfaceNumber - base {
case 0:
ports[index].Role = PortRoleDiagnostic
case 1:
ports[index].Role = PortRoleNMEA
case 2:
ports[index].Role = PortRoleAT
case 3:
ports[index].Role = PortRoleModem
default:
ports[index].Role = PortRoleUnknown
}
}
}
func quecPortRole(interfaceNumber int, name string) PortRole {
// Initial best effort. assignQuectelPortRoles replaces this once every
// interface belonging to the same physical modem has been collected.
switch interfaceNumber {
case 0x02:
case 0x00:
return PortRoleDiagnostic
case 0x03:
case 0x01:
return PortRoleNMEA
case 0x04:
case 0x02:
return PortRoleAT
case 0x05:
case 0x03:
return PortRoleModem
default:
if name == "ttyUSB2" {
@@ -279,9 +295,9 @@ func selectATPort(ports []Port) Port {
for _, port := range ports {
score := 0
switch {
case port.Name == "ttyUSB2":
score = 120
case port.Role == PortRoleAT:
score = 120
case port.Name == "ttyUSB2":
score = 100
case port.InterfaceNumber == 0x04:
score = 90
+48
View File
@@ -123,6 +123,54 @@ func TestSysFSDiscoverySelectsTTYUSB2InQMIInterface00Layout(t *testing.T) {
}
}
func TestSysFSDiscoverySelectsATPortForSecondQMIUSBModem(t *testing.T) {
root := t.TempDir()
sysRoot := filepath.Join(root, "sys")
devRoot := filepath.Join(root, "dev")
usbRoot := filepath.Join(sysRoot, "bus", "usb", "devices")
for _, modem := range []struct {
usbName string
ttys []string
wdm string
}{
{"1-6", []string{"ttyUSB0", "ttyUSB1", "ttyUSB2", "ttyUSB3"}, "cdc-wdm0"},
{"1-5", []string{"ttyUSB4", "ttyUSB5", "ttyUSB6", "ttyUSB7"}, "cdc-wdm1"},
} {
mustWrite(t, filepath.Join(usbRoot, modem.usbName, "idVendor"), "2c7c\n")
mustWrite(t, filepath.Join(usbRoot, modem.usbName, "idProduct"), "0125\n")
for number, tty := range modem.ttys {
interfaceName := modem.usbName + ":1." + strconv.Itoa(number)
mustWrite(t, filepath.Join(usbRoot, interfaceName, "bInterfaceNumber"), fmt.Sprintf("%02x\n", number))
mustMkdir(t, filepath.Join(usbRoot, interfaceName, tty, "tty", tty))
}
mustWrite(t, filepath.Join(usbRoot, modem.usbName+":1.4", "bInterfaceNumber"), "04\n")
mustMkdir(t, filepath.Join(usbRoot, modem.usbName+":1.4", "usbmisc", modem.wdm))
}
candidates, err := NewSysFSDiscoverer(sysRoot, devRoot).Discover(context.Background())
if err != nil {
t.Fatalf("Discover: %v", err)
}
if len(candidates) != 2 {
t.Fatalf("got %d candidates, want 2", len(candidates))
}
for _, candidate := range candidates {
switch filepath.Base(candidate.USBPath) {
case "1-5":
if candidate.ATPort.Name != "ttyUSB6" || candidate.ATPort.Role != PortRoleAT {
t.Fatalf("second modem AT port = %#v, want ttyUSB6", candidate.ATPort)
}
case "1-6":
if candidate.ATPort.Name != "ttyUSB2" || candidate.ATPort.Role != PortRoleAT {
t.Fatalf("first modem AT port = %#v, want ttyUSB2", candidate.ATPort)
}
default:
t.Fatalf("unexpected candidate USB path %q", candidate.USBPath)
}
}
}
func TestSysFSDiscoveryIgnoresNonQuectelUSB(t *testing.T) {
root := t.TempDir()
usbRoot := filepath.Join(root, "sys", "bus", "usb", "devices")
+270
View File
@@ -0,0 +1,270 @@
package server
import (
"context"
"errors"
"net/http"
"strconv"
"strings"
"time"
"vocat/internal/modem"
"vocat/internal/store"
"vocat/internal/vowifi"
)
const maxCallDuration = 10 * time.Minute
func (s *Server) handleCalls(w http.ResponseWriter, r *http.Request, config store.Device, physicalID string) bool {
if !requireMethod(w, r, http.MethodGet) {
return true
}
transport := s.callTransport(config.ID)
if transport == "vowifi" {
controller, ok := s.vowifi.(VoWiFiCallController)
if !ok {
writeError(w, http.StatusNotImplemented, "vowifi_voice_unavailable", "the active VoWiFi IMS session does not expose voice-call signalling")
return true
}
calls, err := controller.Calls(config.ID)
if err != nil {
writeError(w, http.StatusServiceUnavailable, "vowifi_call_failed", err.Error())
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{
"device_id": config.ID, "transport": transport, "calls": calls,
}})
return true
}
response, err := s.devices.ExecuteAT(r.Context(), physicalID, "AT+CLCC")
if err != nil {
s.writeDeviceError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"device_id": config.ID,
"transport": transport,
"calls": parseCLCC(response),
"raw": response.Text(),
},
})
return true
}
func (s *Server) handleCallAction(w http.ResponseWriter, r *http.Request, config store.Device, physicalID, action string) bool {
if !requireMethod(w, r, http.MethodPost) {
return true
}
command := ""
duration := time.Duration(0)
number := ""
callID := ""
switch action {
case "dial":
var request struct {
Number string `json:"number"`
DurationSeconds int `json:"duration_seconds"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
number = strings.TrimSpace(request.Number)
if !validDialNumber(number) {
writeError(w, http.StatusBadRequest, "invalid_number", "phone number is invalid")
return true
}
duration = time.Duration(request.DurationSeconds) * time.Second
if duration < 0 || duration > maxCallDuration {
writeError(w, http.StatusBadRequest, "invalid_duration", "duration_seconds must be 0 (no automatic hang-up) or between 1 and 600")
return true
}
command = "ATD" + number + ";"
case "answer":
var request struct {
CallID string `json:"call_id"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
callID = strings.TrimSpace(request.CallID)
command = "ATA"
case "hangup":
var request struct {
CallID string `json:"call_id"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
callID = strings.TrimSpace(request.CallID)
command = "ATH"
default:
writeError(w, http.StatusNotFound, "not_found", "call action not found")
return true
}
transport := s.callTransport(config.ID)
if transport == "vowifi" {
controller, ok := s.vowifi.(VoWiFiCallController)
if !ok {
writeError(w, http.StatusNotImplemented, "vowifi_voice_unavailable", "the active VoWiFi IMS session does not expose voice-call signalling")
return true
}
var result any
var err error
switch action {
case "dial":
result, err = controller.DialCall(r.Context(), config.ID, number)
case "answer":
callID, err = resolveVoWiFiCallID(controller, config.ID, callID, "ringing")
if err == nil {
result, err = controller.AnswerCall(r.Context(), config.ID, callID)
}
case "hangup":
callID, err = resolveVoWiFiCallID(controller, config.ID, callID, "")
if err == nil {
err = controller.HangupCall(r.Context(), config.ID, callID)
}
}
if err != nil {
writeError(w, http.StatusBadGateway, "vowifi_call_failed", err.Error())
return true
}
if action == "dial" {
if call, ok := result.(vowifi.Call); ok {
callID = call.ID
}
if duration > 0 {
go s.hangupVoWiFiAfter(config.ID, callID, duration)
}
}
s.recordAudit(r.Context(), "admin", "call."+action, "device", config.ID, "success", transport)
writeJSON(w, http.StatusAccepted, map[string]any{"data": map[string]any{
"accepted": true, "action": action, "number": number, "call_id": callID,
"duration_seconds": int(duration / time.Second), "transport": transport, "call": result,
}})
return true
}
operationContext, cancel := context.WithTimeout(r.Context(), 20*time.Second)
response, err := s.devices.ExecuteAT(operationContext, physicalID, command)
cancel()
if err != nil {
s.writeDeviceError(w, err)
return true
}
if !strings.EqualFold(strings.TrimSpace(response.Final), "OK") {
writeError(w, http.StatusBadGateway, "call_rejected", "modem did not accept the call action")
return true
}
if action == "dial" {
if duration > 0 {
go s.hangupAfter(config.ID, physicalID, duration)
}
}
s.recordAudit(r.Context(), "admin", "call."+action, "device", config.ID, "success", transport)
writeJSON(w, http.StatusAccepted, map[string]any{
"data": map[string]any{
"accepted": true, "action": action, "number": number,
"duration_seconds": int(duration / time.Second), "transport": transport,
},
})
return true
}
func resolveVoWiFiCallID(controller VoWiFiCallController, deviceID, id, requiredState string) (string, error) {
if id != "" {
return id, nil
}
calls, err := controller.Calls(deviceID)
if err != nil {
return "", err
}
for _, call := range calls {
if call.State != "ended" && call.State != "failed" && (requiredState == "" || call.State == requiredState) {
return call.ID, nil
}
}
return "", errors.New("no matching active call")
}
func (s *Server) hangupVoWiFiAfter(deviceID, callID string, duration time.Duration) {
timer := time.NewTimer(duration)
defer timer.Stop()
<-timer.C
controller, ok := s.vowifi.(VoWiFiCallController)
if !ok {
return
}
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
if err := controller.HangupCall(ctx, deviceID, callID); err != nil {
s.logger.Warn("automatic VoWiFi call hangup failed", "device_id", deviceID, "call_id", callID, "error", err)
}
}
func (s *Server) callTransport(deviceID string) string {
if s.vowifi != nil {
// Enabled is only the desired card policy. Calls can use IMS only after
// registration has actually completed; otherwise keep using the modem's
// circuit-switched call path instead of routing into an unavailable IMS
// session.
if state, err := s.vowifi.State(deviceID); err == nil && state.IMSReady {
return "vowifi"
}
}
return "cellular"
}
func (s *Server) hangupAfter(deviceID, physicalID string, duration time.Duration) {
timer := time.NewTimer(duration)
defer timer.Stop()
<-timer.C
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
defer cancel()
if _, err := s.devices.ExecuteAT(ctx, physicalID, "ATH"); err != nil {
s.logger.Warn("automatic call hangup failed", "device_id", deviceID, "error", err)
}
}
func validDialNumber(value string) bool {
if len(value) < 2 || len(value) > 32 {
return false
}
for index, character := range value {
if character >= '0' && character <= '9' || (index == 0 && character == '+') || character == '*' || character == '#' {
continue
}
return false
}
return true
}
func parseCLCC(response modem.Response) []map[string]any {
result := make([]map[string]any, 0)
for _, line := range response.Lines {
line = strings.TrimSpace(line)
if !strings.HasPrefix(strings.ToUpper(line), "+CLCC:") {
continue
}
fields := strings.Split(strings.TrimSpace(strings.TrimPrefix(line, "+CLCC:")), ",")
if len(fields) < 5 {
continue
}
integer := func(index int) int {
value, _ := strconv.Atoi(strings.TrimSpace(fields[index]))
return value
}
call := map[string]any{
"index": integer(0), "direction": integer(1), "state": integer(2),
"mode": integer(3), "multiparty": integer(4), "raw": line,
}
if len(fields) > 5 {
call["number"] = strings.Trim(strings.TrimSpace(fields[5]), `"`)
}
result = append(result, call)
}
return result
}
+73
View File
@@ -0,0 +1,73 @@
package server
import (
"context"
"testing"
"vocat/internal/modem"
"vocat/internal/vowifi"
)
func TestParseCLCC(t *testing.T) {
calls := parseCLCC(modem.Response{Lines: []string{
`+CLCC: 1,1,4,0,0,"+447700900000",145`,
`+CLCC: 2,0,0,0,0,"12345",129`,
}})
if len(calls) != 2 || calls[0]["number"] != "+447700900000" || calls[1]["state"] != 0 {
t.Fatalf("parseCLCC = %#v", calls)
}
}
func TestValidDialNumber(t *testing.T) {
for _, value := range []string{"+447700900000", "12345", "*100#"} {
if !validDialNumber(value) {
t.Errorf("validDialNumber(%q) = false", value)
}
}
for _, value := range []string{"", "+", "12;ATH", "12 34", "abc"} {
if validDialNumber(value) {
t.Errorf("validDialNumber(%q) = true", value)
}
}
}
func TestCallTransportRequiresIMSReady(t *testing.T) {
controller := &fakeVoWiFiController{state: vowifi.State{Enabled: true}}
server := &Server{vowifi: controller}
if got := server.callTransport("ec20"); got != "cellular" {
t.Fatalf("callTransport before IMS registration = %q, want cellular", got)
}
controller.state.IMSReady = true
if got := server.callTransport("ec20"); got != "vowifi" {
t.Fatalf("callTransport with IMS ready = %q, want vowifi", got)
}
}
func TestResolveVoWiFiCallIDIgnoresTerminalCalls(t *testing.T) {
controller := &fakeCallController{calls: []vowifi.Call{
{ID: "failed", State: "failed"},
{ID: "active", State: "active"},
}}
got, err := resolveVoWiFiCallID(controller, "ec20", "", "")
if err != nil || got != "active" {
t.Fatalf("resolveVoWiFiCallID() = %q, %v; want active", got, err)
}
}
type fakeCallController struct {
calls []vowifi.Call
}
func (controller *fakeCallController) Calls(string) ([]vowifi.Call, error) {
return controller.calls, nil
}
func (*fakeCallController) DialCall(context.Context, string, string) (vowifi.Call, error) {
return vowifi.Call{}, nil
}
func (*fakeCallController) AnswerCall(context.Context, string, string) (vowifi.Call, error) {
return vowifi.Call{}, nil
}
func (*fakeCallController) HangupCall(context.Context, string, string) error { return nil }
+99
View File
@@ -0,0 +1,99 @@
package server
import (
"context"
"encoding/binary"
"errors"
"io"
"net/http"
"strings"
"github.com/coder/websocket"
"vocat/internal/store"
)
const maxCallMediaMessage = 16 << 10
// handleCallMedia upgrades an authenticated same-origin request to a binary
// PCM bridge. Each WebSocket message contains little-endian signed 16-bit,
// 8 kHz, mono samples. RTP and codec details remain inside the IMS provider.
func (s *Server) handleCallMedia(w http.ResponseWriter, r *http.Request, config store.Device) bool {
if !requireMethod(w, r, http.MethodGet) {
return true
}
if s.callTransport(config.ID) != "vowifi" {
writeError(w, http.StatusNotImplemented, "call_media_unavailable", "browser audio is only available for an active VoWiFi IMS call")
return true
}
callID := strings.TrimSpace(r.URL.Query().Get("call_id"))
if callID == "" || len(callID) > 256 {
writeError(w, http.StatusBadRequest, "invalid_call_id", "call_id is required")
return true
}
controller, ok := s.vowifi.(VoWiFiCallMediaController)
if !ok {
writeError(w, http.StatusNotImplemented, "call_media_unavailable", "the active IMS session does not expose RTP media")
return true
}
media, err := controller.CallMedia(r.Context(), config.ID, callID)
if err != nil {
writeError(w, http.StatusConflict, "call_media_unavailable", err.Error())
return true
}
connection, err := websocket.Accept(w, r, &websocket.AcceptOptions{
CompressionMode: websocket.CompressionDisabled,
})
if err != nil {
return true
}
connection.SetReadLimit(maxCallMediaMessage)
ctx, cancel := context.WithCancel(r.Context())
defer cancel()
defer connection.Close(websocket.StatusNormalClosure, "call media closed")
downlink := make(chan error, 1)
go func() {
defer cancel()
for {
samples, readErr := media.ReadPCM(ctx)
if readErr != nil {
downlink <- readErr
return
}
payload := make([]byte, len(samples)*2)
for index, sample := range samples {
binary.LittleEndian.PutUint16(payload[index*2:], uint16(sample))
}
if writeErr := connection.Write(ctx, websocket.MessageBinary, payload); writeErr != nil {
downlink <- writeErr
return
}
}
}()
for {
select {
case err := <-downlink:
if !errors.Is(err, context.Canceled) && !errors.Is(err, io.EOF) {
s.logger.Debug("call media downlink closed", "device_id", config.ID, "call_id", callID, "error", err)
}
return true
default:
}
messageType, payload, readErr := connection.Read(ctx)
if readErr != nil {
return true
}
if messageType != websocket.MessageBinary || len(payload) == 0 || len(payload)%2 != 0 {
continue
}
samples := make([]int16, len(payload)/2)
for index := range samples {
samples[index] = int16(binary.LittleEndian.Uint16(payload[index*2:]))
}
if err := media.WritePCM(samples); err != nil {
return true
}
}
}
+43
View File
@@ -0,0 +1,43 @@
package server
import (
"net/http"
"vocat/internal/developer"
)
func (s *Server) handleDeveloperSettings(w http.ResponseWriter, r *http.Request) {
if !s.developerEnabled {
writeError(w, http.StatusNotFound, "not_found", "resource not found")
return
}
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{
"device_limit": developer.DeviceLimit(r.Context(), s.store, true),
"default_device_limit": developer.DefaultDeviceLimit,
"max_device_limit": developer.MaxDeviceLimit,
}})
case http.MethodPut:
var request struct {
DeviceLimit int `json:"device_limit"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return
}
if err := developer.SetDeviceLimit(r.Context(), s.store, request.DeviceLimit); err != nil {
writeError(w, http.StatusBadRequest, "invalid_device_limit", err.Error())
return
}
s.recordAudit(r.Context(), "admin", "settings.developer.device_limit", "settings", "developer", "success", "device limit updated")
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{
"device_limit": request.DeviceLimit,
"default_device_limit": developer.DefaultDeviceLimit,
"max_device_limit": developer.MaxDeviceLimit,
}})
default:
w.Header().Set("Allow", "GET, PUT")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
}
@@ -0,0 +1,22 @@
package server
import (
"net/http"
"net/http/httptest"
"testing"
)
func TestDeveloperOnlySettingsAreHiddenWhenModeIsOff(t *testing.T) {
server := &Server{developerEnabled: false}
for _, handler := range []func(http.ResponseWriter, *http.Request){
server.handleDeveloperSettings,
server.handleHTTPSSettings,
server.handleHTTPSCertificate,
} {
response := httptest.NewRecorder()
handler(response, httptest.NewRequest(http.MethodGet, "/api/settings/developer", nil))
if response.Code != http.StatusNotFound {
t.Fatalf("developer-only endpoint status = %d, want 404", response.Code)
}
}
}
+351 -90
View File
@@ -11,6 +11,7 @@ import (
"strings"
"time"
"vocat/internal/developer"
"vocat/internal/device"
"vocat/internal/i18n"
"vocat/internal/modem"
@@ -38,6 +39,7 @@ type DeviceController interface {
SetUSBNetModeByPort(context.Context, string, int) (device.USBNetMode, error)
OperatorSelection(context.Context, string) (device.OperatorSelection, error)
SetOperatorSelection(context.Context, string, bool, string, *int) (device.OperatorSelection, error)
ReRegisterOperator(context.Context, string) (device.OperatorSelection, error)
ScanOperators(context.Context, string) (device.OperatorScanResult, error)
SendSMS(context.Context, string, string, string) (device.SMSSendResult, error)
ListSMS(context.Context, string) ([]device.SMSMessage, error)
@@ -56,6 +58,7 @@ type DeviceController interface {
type deviceConfigPayload struct {
ID string `json:"id"`
Name string `json:"name"`
DeviceType string `json:"device_type"`
Interface string `json:"interface"`
ControlDevice string `json:"control_device"`
ATPort string `json:"at_port"`
@@ -86,6 +89,7 @@ func (payload deviceConfigPayload) toStoreDevice() store.Device {
return store.Device{
ID: strings.TrimSpace(payload.ID),
Name: name,
DeviceType: store.NormalizeDeviceType(payload.DeviceType),
Interface: strings.TrimSpace(payload.Interface),
ControlDevice: strings.TrimSpace(payload.ControlDevice),
ATPort: strings.TrimSpace(payload.ATPort),
@@ -170,15 +174,13 @@ func splitAPIPath(value string) []string {
return result
}
// maxDeviceLimit 是设备数量的软上限:达到上限后禁止再添加新设备。
const maxDeviceLimit = 5
func (s *Server) handleDevices(w http.ResponseWriter, r *http.Request) bool {
deviceLimit := developer.DeviceLimit(r.Context(), s.store, s.developerEnabled)
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"device_limit": maxDeviceLimit,
"device_limit": deviceLimit,
"devices": s.deviceSummaries(),
},
})
@@ -203,6 +205,10 @@ func (s *Server) handleDevices(w http.ResponseWriter, r *http.Request) bool {
writeError(w, http.StatusBadRequest, "invalid_device_id", "device ID must use 1-64 letters, digits, dots, underscores, or hyphens")
return true
}
if strings.TrimSpace(payload.DeviceType) == "" || store.NormalizeDeviceType(payload.DeviceType) == "" {
writeError(w, http.StatusBadRequest, "invalid_device_type", "select a supported device type")
return true
}
if _, err := s.store.Device(r.Context(), payload.ID); err == nil {
writeError(w, http.StatusConflict, "device_exists", "a device with this ID already exists")
return true
@@ -215,8 +221,8 @@ func (s *Server) handleDevices(w http.ResponseWriter, r *http.Request) bool {
s.writeStoreError(w, err)
return true
}
if len(configured) >= maxDeviceLimit {
writeError(w, http.StatusConflict, "device_limit_reached", i18n.Tf("设备数量已达上限,最多只能添加 %d 台设备", maxDeviceLimit))
if len(configured) >= deviceLimit {
writeError(w, http.StatusConflict, "device_limit_reached", i18n.Tf("设备数量已达上限,最多只能添加 %d 台设备", deviceLimit))
return true
}
devices, err := s.devices.Discover(r.Context())
@@ -230,6 +236,9 @@ func (s *Server) handleDevices(w http.ResponseWriter, r *http.Request) bool {
return true
}
config := payload.toStoreDevice()
if !s.developerActive(r.Context()) {
config.NetworkEnabled = false
}
fillConfigFromPhysical(&config, *selected)
if err := s.store.UpsertDevice(r.Context(), config); err != nil {
s.writeStoreError(w, err)
@@ -251,18 +260,33 @@ func (s *Server) handleDevices(w http.ResponseWriter, r *http.Request) bool {
}
func findDiscoveredDevice(devices []device.Device, config deviceConfigPayload) *device.Device {
if config.ModemIMEI != "" {
for index := range devices {
if devices[index].Snapshot != nil && devices[index].Snapshot.IMEI == config.ModemIMEI {
return &devices[index]
}
}
}
if config.USBPath != "" {
for index := range devices {
if devices[index].Candidate.USBPath == config.USBPath {
return &devices[index]
}
}
}
if config.ControlDevice != "" {
for index := range devices {
if devices[index].Candidate.QMIControl == config.ControlDevice {
return &devices[index]
}
}
}
for index := range devices {
candidate := devices[index].Candidate
if config.ATPort != "" &&
(candidate.ATPort.Path == config.ATPort || candidate.ATPort.OpenPath() == config.ATPort) {
return &devices[index]
}
if config.ControlDevice != "" && candidate.QMIControl == config.ControlDevice {
return &devices[index]
}
if config.USBPath != "" && candidate.USBPath == config.USBPath {
return &devices[index]
}
}
return nil
}
@@ -387,6 +411,9 @@ func (s *Server) handleDevicePath(
return true
}
next := payload.toStoreDevice()
if !s.developerActive(r.Context()) {
next.NetworkEnabled = false
}
next.ID = id
next.CreatedAt = config.CreatedAt
if next.Name == id && strings.TrimSpace(payload.Name) == "" {
@@ -470,12 +497,27 @@ func (s *Server) handleDevicePath(
s.writeDeviceError(w, err)
return true
}
s.clearPublicIP(config.ID)
writeJSON(w, http.StatusAccepted, map[string]any{"data": map[string]any{"status": "rebooting"}})
case "flight-mode":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
return s.handleFlightMode(w, r, physicalID)
case "network":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
return s.handleCellularData(w, r, config, physicalID)
case "network/public-ip":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
iccid := ""
if entry.Snapshot != nil {
iccid = entry.Snapshot.ICCID
}
return s.handleCellularPublicIP(w, r, config, iccid)
case "usbnet-mode":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
@@ -486,6 +528,11 @@ func (s *Server) handleDevicePath(
return true
}
return s.handleOperatorSelection(w, r, physicalID)
case "operator_selection/reregister":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
return s.handleOperatorReRegister(w, r, physicalID)
case "operator_selection/scan":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
@@ -502,6 +549,21 @@ func (s *Server) handleDevicePath(
return s.handleVoWiFiReconnect(w, r, config, physicalPresent)
case "vowifi/e911/websheet":
return s.handleE911Websheet(w, r, config)
case "calls":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
return s.handleCalls(w, r, config, physicalID)
case "calls/dial", "calls/answer", "calls/hangup":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
return s.handleCallAction(w, r, config, physicalID, tail[1])
case "calls/media":
if !s.requirePhysicalDevice(w, physicalPresent) {
return true
}
return s.handleCallMedia(w, r, config)
default:
return false
}
@@ -642,6 +704,21 @@ func (s *Server) handleOperatorSelection(w http.ResponseWriter, r *http.Request,
return true
}
func (s *Server) handleOperatorReRegister(w http.ResponseWriter, r *http.Request, physicalID string) bool {
if !requireMethod(w, r, http.MethodPost) {
return true
}
controller := http.NewResponseController(w)
_ = controller.SetWriteDeadline(time.Time{})
result, err := s.devices.ReRegisterOperator(r.Context(), physicalID)
if err != nil {
s.writeDeviceError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": operatorSelectionWire(result)})
return true
}
func (s *Server) handleVoWiFiEnabled(
w http.ResponseWriter,
r *http.Request,
@@ -666,6 +743,10 @@ func (s *Server) handleVoWiFiEnabled(
writeError(w, http.StatusServiceUnavailable, "physical_device_missing", "the configured modem is not present on this Linux host")
return true
}
if request.Enabled && config.NetworkEnabled {
writeError(w, http.StatusConflict, "cellular_data_active", "disable roaming data before enabling VoWiFi")
return true
}
if request.Enabled {
entry, _, _ := s.physicalForConfig(config)
imsi := snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMSI })
@@ -918,10 +999,111 @@ func (s *Server) handleFlightMode(w http.ResponseWriter, r *http.Request, id str
s.writeDeviceError(w, err)
return true
}
// Unlike VoWiFi, CFUN airplane state is not represented in the device row.
// Persist it against the live ICCID so a restart can distinguish an
// intentional airplane policy from an interrupted VoWiFi teardown.
if entry, getErr := s.devices.Get(id); getErr == nil && entry.Snapshot != nil {
iccid := strings.TrimSpace(entry.Snapshot.ICCID)
if iccid != "" {
policy, policyErr := s.store.CardPolicy(r.Context(), iccid)
if errors.Is(policyErr, store.ErrNotFound) {
policy = store.CardPolicy{ICCID: iccid, IPVersion: "IPV4V6"}
policyErr = nil
}
if policyErr != nil {
s.writeStoreError(w, policyErr)
return true
}
policy.AirplaneEnabled = request.Enabled
if request.Enabled {
policy.VoWiFiEnabled = false
}
policy.Source = "manual"
if err := s.store.UpsertCardPolicy(r.Context(), policy); err != nil {
s.writeStoreError(w, err)
return true
}
}
}
writeJSON(w, http.StatusOK, map[string]any{"data": result})
return true
}
func (s *Server) handleCellularData(
w http.ResponseWriter,
r *http.Request,
config store.Device,
physicalID string,
) bool {
if !s.developerActive(r.Context()) {
writeError(w, http.StatusForbidden, "developer_mode_required", "roaming data is available only in developer mode")
return true
}
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{
"enabled": config.NetworkEnabled,
"interface": config.Interface,
"apn": config.APN,
"export_proxy_only": true,
}})
case http.MethodPatch, http.MethodPut:
var request struct {
Enabled bool `json:"enabled"`
APN string `json:"apn"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
if request.Enabled && config.VoWiFiEnabled {
writeError(w, http.StatusConflict, "vowifi_owns_radio", "disable VoWiFi before enabling cellular roaming data")
return true
}
if !request.Enabled && s.exportProxy != nil {
if _, active := s.exportProxy.EnabledConfigForDevice(config.ID); active {
writeError(w, http.StatusConflict, "export_proxy_active", i18n.T("请先禁用该设备已绑定的导出代理,再关闭漫游数据"))
return true
}
}
apn := strings.TrimSpace(request.APN)
if apn == "" {
apn = strings.TrimSpace(config.APN)
}
controller := http.NewResponseController(w)
_ = controller.SetWriteDeadline(time.Time{})
result, err := s.devices.SetNetwork(r.Context(), physicalID, device.NetworkRequest{
Enabled: request.Enabled, APN: apn, IPVersion: "IPV4V6",
})
if err != nil {
s.writeDeviceError(w, err)
return true
}
previous := config.NetworkEnabled
config.NetworkEnabled = request.Enabled
if apn != "" {
config.APN = apn
}
if err := s.store.UpsertDevice(r.Context(), config); err != nil {
rollbackContext, cancel := context.WithTimeout(context.Background(), 20*time.Second)
_, _ = s.devices.SetNetwork(rollbackContext, physicalID, device.NetworkRequest{
Enabled: previous, APN: config.APN, IPVersion: "IPV4V6",
})
cancel()
s.writeStoreError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{
"enabled": result.Enabled, "interface": result.Interface,
"backend": result.Backend, "export_proxy_only": true,
}})
default:
w.Header().Set("Allow", "GET, PATCH, PUT")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
return true
}
func (s *Server) requirePhysicalDevice(w http.ResponseWriter, present bool) bool {
if s.devices == nil {
writeError(w, http.StatusServiceUnavailable, "device_manager_unavailable", "device manager is unavailable")
@@ -1006,6 +1188,7 @@ func (s *Server) dashboardDevices() []map[string]any {
result = append(result, map[string]any{
"id": entry["id"],
"name": entry["name"],
"device_type": entry["device_type"],
"interface": entry["interface"],
"proxy_port": entry["proxy_port"],
"public_ip": entry["public_ip"],
@@ -1016,7 +1199,7 @@ func (s *Server) dashboardDevices() []map[string]any {
"network_duplex": modemStatus["network_duplex"],
"vowifi_active": vowifiActive,
"vowifi_runtime": runtime,
"network_connected": false,
"network_connected": entry["network_connected"],
"model": modemStatus["model"],
})
}
@@ -1043,6 +1226,14 @@ func physicalMatchesConfig(entry device.Device, config store.Device) bool {
if entry.ID == config.ID {
return true
}
if config.ModemIMEI != "" && entry.Snapshot != nil && entry.Snapshot.IMEI != "" {
return config.ModemIMEI == entry.Snapshot.IMEI
}
if config.USBPath != "" && candidate.USBPath != "" {
return config.USBPath == candidate.USBPath
}
// Control and serial device nodes are allocation-order dependent. They are
// only legacy fallbacks when no physical USB path or readable IMEI exists.
if config.ATPort != "" &&
(config.ATPort == candidate.ATPort.Path || config.ATPort == candidate.ATPort.OpenPath()) {
return true
@@ -1050,12 +1241,7 @@ func physicalMatchesConfig(entry device.Device, config store.Device) bool {
if config.ControlDevice != "" && config.ControlDevice == candidate.QMIControl {
return true
}
if config.USBPath != "" && config.USBPath == candidate.USBPath {
return true
}
return config.ModemIMEI != "" &&
entry.Snapshot != nil &&
config.ModemIMEI == entry.Snapshot.IMEI
return false
}
func (s *Server) configuredDeviceSummary(
@@ -1070,24 +1256,50 @@ func (s *Server) configuredDeviceSummary(
}
result["id"] = config.ID
result["name"] = config.Name
result["device_type"] = store.NormalizeDeviceType(config.DeviceType)
result["interface"] = config.Interface
result["proxy_port"] = config.ProxyPort
result["esim_transport"] = config.ESIMTransport
result["sms_enabled"] = config.SMSEnabled
result["network_enabled"] = false
result["network_enabled"] = config.NetworkEnabled
result["developer_enabled"] = s.developerActive(context.Background())
result["network_connected"] = config.NetworkEnabled
result["data_connected"] = config.NetworkEnabled
result["vowifi_enabled"] = config.VoWiFiEnabled
if runtime, err := s.store.VoWiFiRuntime(context.Background(), config.ID); err == nil {
runtimeResponse := storedVoWiFiRuntime(runtime)
currentICCID := ""
var currentSnapshot *device.Snapshot
if entry != nil {
currentSnapshot = entry.Snapshot
if entry.Snapshot != nil {
currentICCID = strings.TrimSpace(entry.Snapshot.ICCID)
}
}
runtimeMatchesCard := currentICCID == "" || runtime.ICCID == "" ||
strings.EqualFold(currentICCID, strings.TrimSpace(runtime.ICCID))
var runtimeResponse map[string]any
if runtimeMatchesCard {
runtimeResponse = storedVoWiFiRuntime(runtime)
} else {
// The saved IMS session belongs to a different eSIM profile. Never
// project its registration or number onto the currently selected SIM.
runtimeResponse = idleVoWiFiRuntime(config.ID, currentSnapshot)
}
result["vowifi_runtime"] = runtimeResponse
result["vowifi_active"] = runtime.TunnelReady
if runtime.LocalPhone != "" {
// The SIM panel reads the top-level local_phone; keep modem.phone_number
// in sync for the summary/overview consumers that read it there.
result["local_phone"] = runtime.LocalPhone
result["phone_number_source"] = runtime.PhoneNumberSource
if modemStatus, ok := result["modem"].(map[string]any); ok {
modemStatus["phone_number"] = runtime.LocalPhone
modemStatus["phone_number_source"] = runtime.PhoneNumberSource
result["vowifi_active"] = config.VoWiFiEnabled && runtimeMatchesCard && runtime.TunnelReady
}
// Numbers are SIM-owned data. Resolve the association by the live ICCID
// instead of reusing the last VoWiFi runtime attached to this device ID.
if entry != nil && entry.Snapshot != nil {
currentICCID := strings.TrimSpace(entry.Snapshot.ICCID)
if currentICCID != "" {
if association, err := s.store.PhoneAssociation(context.Background(), currentICCID); err == nil {
result["local_phone"] = association.Number
result["phone_number_source"] = association.Source
if modemStatus, ok := result["modem"].(map[string]any); ok {
modemStatus["phone_number"] = association.Number
modemStatus["phone_number_source"] = association.Source
}
}
}
}
@@ -1099,6 +1311,7 @@ func (s *Server) configuredDeviceOverview(
entry device.Device,
present bool,
) map[string]any {
developerActive := s.developerActive(context.Background())
var physical *device.Device
if present {
physical = &entry
@@ -1113,12 +1326,34 @@ func (s *Server) configuredDeviceOverview(
result["control_device"] = config.ControlDevice
result["esim_transport"] = config.ESIMTransport
result["sms_enabled"] = config.SMSEnabled
result["network_enabled"] = false
result["network_enabled"] = developerActive && config.NetworkEnabled
result["vowifi_enabled"] = config.VoWiFiEnabled
result["radio_live_ok"] = present && entry.Snapshot != nil && entry.Snapshot.Responsive
result["traffic"] = map[string]string{}
result["traffic_raw"] = map[string]int64{}
result["traffic_meta"] = map[string]any{}
// Live network state: on-demand sample of the cellular interface counters,
// kept warm by the 2s overview SSE cadence. Only meaningful when the modem
// data path is enabled and an interface is configured.
if developerActive && config.NetworkEnabled && strings.TrimSpace(config.Interface) != "" {
live := s.netTraffic.sample(config.ID, config.Interface, time.Now())
result["private_ip"] = live.ipv4
result["traffic"] = map[string]string{
"rx": formatLiveBytes(float64(live.minuteRx)),
"tx": formatLiveBytes(float64(live.minuteTx)),
"rate": formatLiveBytes(live.rxRate) + "/s",
"rate_tx": formatLiveBytes(live.txRate) + "/s",
}
result["traffic_raw"] = map[string]int64{
"rx": live.minuteRx,
"tx": live.minuteTx,
"rate": int64(live.rxRate),
"rate_tx": int64(live.txRate),
}
result["traffic_meta"] = map[string]any{"status": live.status}
} else {
result["traffic"] = map[string]string{}
result["traffic_raw"] = map[string]int64{}
result["traffic_meta"] = map[string]any{}
}
return result
}
@@ -1139,7 +1374,7 @@ func (s *Server) configuredDeviceStatus(
result := map[string]any{
"healthy": summary["healthy"],
"public_ip": summary["public_ip"],
"network_connected": false,
"network_connected": config.NetworkEnabled,
"modem": summary["modem"],
"vowifi": summary["vowifi_runtime"],
"sim_service_table": map[string]any{},
@@ -1222,7 +1457,7 @@ func deviceSummary(entry device.Device) map[string]any {
"physical_present": entry.Discovered,
"worker_running": entry.Discovered,
"data_connected": false,
"radio_registered": snapshot != nil && snapshot.OperatorName != "",
"radio_registered": snapshot != nil && (snapshot.RegistrationStatus == 1 || snapshot.RegistrationStatus == 5),
"lifecycle_phase": lifecyclePhase(entry),
"lifecycle_reason": entry.LastError,
"public_ip": "",
@@ -1279,6 +1514,7 @@ func storedDeviceConfig(config store.Device) map[string]any {
return map[string]any{
"id": config.ID,
"name": config.Name,
"device_type": store.NormalizeDeviceType(config.DeviceType),
"interface": config.Interface,
"control_device": config.ControlDevice,
"at_port": config.ATPort,
@@ -1296,7 +1532,7 @@ func storedDeviceConfig(config store.Device) map[string]any {
"qmi_use_proxy": config.QMIUseProxy,
"qmi_proxy_path": config.QMIProxyPath,
"qmi_proxy_executable": config.QMIProxyExecutable,
"network_enabled": false,
"network_enabled": config.NetworkEnabled,
"sms_enabled": config.SMSEnabled,
"vowifi_enabled": config.VoWiFiEnabled,
}
@@ -1330,61 +1566,64 @@ func fillConfigFromPhysical(config *store.Device, entry device.Device) {
func modemSummary(snapshot *device.Snapshot, phone string, phoneSource string) map[string]any {
if snapshot == nil {
return map[string]any{
"operator": "",
"native_mcc": "",
"native_mnc": "",
"card_mcc": "",
"card_mnc": "",
"card_country": "",
"service_blocked": false,
"blocked_reason": "",
"network_mode": "",
"radio_band": "",
"radio_channel": 0,
"signal_dbm": 0,
"signal_sinr": 0,
"imei": "",
"iccid": "",
"reg_status": 0,
"reg_status_text": "not refreshed",
"sim_inserted": false,
"phone_number": phone,
"phone_number_source": phoneSource,
"model": "",
"operator": "",
"native_mcc": "",
"native_mnc": "",
"operator_country_code": "",
"card_mcc": "",
"card_mnc": "",
"card_country": "",
"service_blocked": false,
"blocked_reason": "",
"network_mode": "",
"radio_band": "",
"radio_channel": 0,
"signal_dbm": 0,
"signal_sinr": 0,
"imei": "",
"iccid": "",
"reg_status": 0,
"reg_status_text": "not refreshed",
"sim_inserted": false,
"phone_number": phone,
"phone_number_source": phoneSource,
"model": "",
}
}
mcc, mnc := splitPLMN(snapshot.OperatorCode)
_, operatorCountryCode, _ := device.CarrierForPLMN(snapshot.OperatorCode)
cardMCC, cardMNC := device.CardMCCMNC(snapshot.IMSI)
blockedReason := device.RegionBlockReason(snapshot.IMSI)
return map[string]any{
"operator": snapshot.OperatorName,
"native_mcc": mcc,
"native_mnc": mnc,
"card_mcc": cardMCC,
"card_mnc": cardMNC,
"card_country": countryNameForMCC(cardMCC),
"service_blocked": blockedReason != "",
"blocked_reason": blockedReason,
"network_mode": snapshot.AccessTech,
"network_duplex": "",
"radio_band": snapshot.Band,
"radio_channel": parseDecimal(snapshot.Channel),
"signal_dbm": pointerInt(snapshot.RSSIDBm),
"signal_rsrp": pointerInt(snapshot.RSRP),
"signal_rsrq": pointerInt(snapshot.RSRQ),
"signal_sinr": pointerInt(snapshot.SINR),
"imei": snapshot.IMEI,
"iccid": snapshot.ICCID,
"imsi": snapshot.IMSI,
"firmware": snapshot.Firmware,
"model": snapshot.Model,
"reg_status": boolInt(snapshot.OperatorName != ""),
"reg_status_text": registrationText(snapshot),
"ps_attached": false,
"sim_inserted": snapshot.SIMStatus != "",
"operating_mode": snapshot.OperatingMode,
"phone_number": phone,
"phone_number_source": phoneSource,
"operator": snapshot.OperatorName,
"native_mcc": mcc,
"native_mnc": mnc,
"operator_country_code": operatorCountryCode,
"card_mcc": cardMCC,
"card_mnc": cardMNC,
"card_country": countryNameForMCC(cardMCC),
"service_blocked": blockedReason != "",
"blocked_reason": blockedReason,
"network_mode": snapshot.AccessTech,
"network_duplex": "",
"radio_band": snapshot.Band,
"radio_channel": parseDecimal(snapshot.Channel),
"signal_dbm": pointerInt(snapshot.RSSIDBm),
"signal_rsrp": pointerInt(snapshot.RSRP),
"signal_rsrq": pointerInt(snapshot.RSRQ),
"signal_sinr": pointerInt(snapshot.SINR),
"imei": snapshot.IMEI,
"iccid": snapshot.ICCID,
"imsi": snapshot.IMSI,
"firmware": snapshot.Firmware,
"model": snapshot.Model,
"reg_status": snapshot.RegistrationStatus,
"reg_status_text": registrationText(snapshot),
"ps_attached": snapshot.PSAttached,
"sim_inserted": snapshot.SIMStatus != "",
"operating_mode": snapshot.OperatingMode,
"phone_number": phone,
"phone_number_source": phoneSource,
}
}
@@ -1457,17 +1696,39 @@ func lifecyclePhase(entry device.Device) string {
}
func registrationLabel(snapshot *device.Snapshot) string {
if snapshot == nil || snapshot.OperatorName == "" {
if snapshot == nil {
return "unknown"
}
switch snapshot.RegistrationStatus {
case 1, 5:
return "registered"
case 2:
return "searching"
case 3:
return "denied"
default:
return "unknown"
}
return "registered"
}
func registrationText(snapshot *device.Snapshot) string {
if snapshot.OperatorName != "" {
if snapshot == nil {
return "unknown"
}
switch snapshot.RegistrationStatus {
case 1:
return "registered"
case 5:
return "registered (roaming)"
case 2:
return "searching"
case 3:
return "registration denied"
case 0:
return "not registered"
default:
return "unknown"
}
return "unknown"
}
func splitPLMN(value string) (string, string) {
+21 -1
View File
@@ -2,6 +2,7 @@ package server
import (
"encoding/json"
"errors"
"fmt"
"net/http"
"time"
@@ -10,6 +11,10 @@ import (
"vocat/internal/store"
)
// overviewStreamInterval is the cadence at which the overview SSE stream pushes
// a fresh snapshot. It is a package var so tests can shorten it.
var overviewStreamInterval = 2 * time.Second
// beginSSE prepares a response for Server-Sent Events and returns its response
// controller for explicit flushes.
func beginSSE(w http.ResponseWriter) *http.ResponseController {
@@ -50,13 +55,27 @@ func (s *Server) handleOverviewStream(
if err := writeSSEEvent(w, controller, "connected", map[string]any{}); err != nil {
return true
}
ticker := time.NewTicker(2 * time.Second)
ticker := time.NewTicker(overviewStreamInterval)
defer ticker.Stop()
for {
select {
case <-r.Context().Done():
return true
case <-ticker.C:
// The config passed in was read when the stream opened. Re-read it on
// every tick so edits made while watching (roaming data, APN, VoWiFi,
// name…) take effect; otherwise the stream keeps replaying the stale
// snapshot and the UI flaps between SSE-old and REST-new values.
fresh, err := s.store.Device(r.Context(), config.ID)
if err != nil {
if errors.Is(err, store.ErrNotFound) {
// The device was deleted while streaming; end the stream.
return true
}
// Transient store hiccup: keep the last known config for this tick.
} else {
config = fresh
}
currentEntry, _, present := s.physicalForConfig(config)
overview := s.configuredDeviceOverview(config, currentEntry, present)
if err := writeSSEEvent(w, controller, "overview", overview); err != nil {
@@ -81,6 +100,7 @@ func operatorCandidateWire(op device.ScannedOperator) map[string]any {
"operatorName": op.Name,
"shortName": op.Short,
"plmn": op.Numeric,
"countryCode": op.Country,
"rats": rats,
"includesPcsDigit": false,
}
+356
View File
@@ -1,15 +1,23 @@
package server
import (
"bufio"
"context"
"encoding/json"
"errors"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"vocat/internal/developer"
"vocat/internal/device"
"vocat/internal/exportproxy"
"vocat/internal/modem"
"vocat/internal/store"
"vocat/internal/update"
)
func decodeData(t *testing.T, recorder *httptest.ResponseRecorder) map[string]any {
@@ -40,6 +48,50 @@ func TestAttachSingleEUICCIdentityFillsProfileGroupMetadataKey(t *testing.T) {
}
}
func TestPhysicalMatchesConfigRejectsDuplicateAndroidSerialAlias(t *testing.T) {
config := store.Device{
ID: "EC20",
ATPort: "/dev/serial/by-id/usb-Android_Android-if02-port0",
USBPath: "/sys/bus/usb/devices/1-6",
ModemIMEI: "111111111111111",
}
newModem := device.Device{
ID: "quectel-0125-1-5",
Candidate: modem.Candidate{
USBPath: "/sys/bus/usb/devices/1-5",
ATPort: modem.Port{
Path: "/dev/ttyUSB6",
StablePath: config.ATPort,
},
},
Snapshot: &device.Snapshot{IMEI: "222222222222222"},
}
if physicalMatchesConfig(newModem, config) {
t.Fatal("different modem matched through a duplicated Android by-id alias")
}
movedOriginal := newModem
movedOriginal.Snapshot = &device.Snapshot{IMEI: config.ModemIMEI}
if !physicalMatchesConfig(movedOriginal, config) {
t.Fatal("same IMEI should follow the modem to a different USB port")
}
}
func TestFindDiscoveredDevicePrefersPhysicalIdentityOverSerialAlias(t *testing.T) {
alias := "/dev/serial/by-id/usb-Android_Android-if02-port0"
devices := []device.Device{
{ID: "old", Candidate: modem.Candidate{USBPath: "/sys/bus/usb/devices/1-6", ATPort: modem.Port{StablePath: alias}}},
{ID: "new", Candidate: modem.Candidate{USBPath: "/sys/bus/usb/devices/1-5", ATPort: modem.Port{StablePath: alias}}},
}
selected := findDiscoveredDevice(devices, deviceConfigPayload{
USBPath: "/sys/bus/usb/devices/1-5",
ATPort: alias,
})
if selected == nil || selected.ID != "new" {
t.Fatalf("selected = %#v, want new physical USB device", selected)
}
}
func TestHandleOperatorScanReturnsOperators(t *testing.T) {
server := &Server{
logger: regionTestLogger(),
@@ -278,6 +330,83 @@ func TestHandleUpdateApplyIsSafeNoop(t *testing.T) {
}
}
func TestHandleUpdateCheckUsesTrustedRepository(t *testing.T) {
server := &Server{
logger: regionTestLogger(),
updateRepository: update.DefaultRepository,
updateCheck: func(_ context.Context, repo, token, current string) (update.CheckResult, error) {
if repo != update.DefaultRepository || token != "token" || current == "" {
t.Fatalf("check arguments = %q, %q, %q", repo, token, current)
}
return update.CheckResult{
Available: true,
Current: current,
Latest: "9.9.9",
ReleaseNotes: "release notes",
}, nil
},
updateToken: "token",
}
recorder := httptest.NewRecorder()
server.handleUpdateCheck(recorder, httptest.NewRequest(http.MethodGet, "/check", nil))
if recorder.Code != http.StatusOK {
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body)
}
data := decodeData(t, recorder)
if data["available"] != true || data["version"] != "9.9.9" || data["repository"] != update.DefaultRepository {
t.Fatalf("check data = %#v", data)
}
}
func TestHandleUpdateApplyInstallsFromTrustedRepository(t *testing.T) {
database, err := store.Open(context.Background(), ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
if err := database.SetAdmin(context.Background(), "admin", []byte("hash")); err != nil {
t.Fatal(err)
}
tokenHash := []byte("active-session")
if err := database.CreateSession(
context.Background(), 1, tokenHash, []byte("csrf"), time.Now().Add(time.Hour),
); err != nil {
t.Fatal(err)
}
server := &Server{
store: database,
logger: regionTestLogger(),
updateRepository: update.DefaultRepository,
updateApply: func(_ context.Context, _ *slog.Logger, options update.Options, restart bool) (update.CheckResult, error) {
if options.Repo != update.DefaultRepository || restart {
t.Fatalf("apply options = %#v, restart = %v", options, restart)
}
return update.CheckResult{Applied: true, Latest: "9.9.9"}, nil
},
}
recorder := httptest.NewRecorder()
server.handleUpdateApply(recorder, httptest.NewRequest(http.MethodPost, "/apply", nil))
if recorder.Code != http.StatusOK {
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body)
}
data := decodeData(t, recorder)
if data["applied"] != true || data["version"] != "9.9.9" || data["reauthentication_required"] != true {
t.Fatalf("apply data = %#v", data)
}
if _, err := database.SessionByTokenHash(context.Background(), tokenHash); !errors.Is(err, store.ErrNotFound) {
t.Fatalf("session must be revoked after update, got %v", err)
}
expired := map[string]bool{}
for _, cookie := range recorder.Result().Cookies() {
if cookie.MaxAge < 0 {
expired[cookie.Name] = true
}
}
if !expired[sessionCookieName] || !expired[csrfCookieName] {
t.Fatalf("auth cookies were not expired: %#v", recorder.Result().Cookies())
}
}
func TestE911WebsheetFlow(t *testing.T) {
database, err := store.Open(context.Background(), ":memory:")
if err != nil {
@@ -341,3 +470,230 @@ func TestE911WebsheetRejectsBadToken(t *testing.T) {
t.Fatalf("bad token status = %d, want 403", recorder.Code)
}
}
// readSSEEvent reads one Server-Sent-Events frame ("event:"/"data:" lines
// terminated by a blank line) and returns the event name and data payload.
func readSSEEvent(reader *bufio.Reader) (string, []byte, error) {
var event string
var data []byte
for {
line, err := reader.ReadString('\n')
if err != nil {
return "", nil, err
}
line = strings.TrimRight(line, "\r\n")
if line == "" {
if event != "" || data != nil {
return event, data, nil
}
continue
}
if rest, ok := strings.CutPrefix(line, "event: "); ok {
event = rest
} else if rest, ok := strings.CutPrefix(line, "data: "); ok {
data = append(data, rest...)
}
}
}
// awaitOverviewNetworkEnabled reads overview SSE events until one reports the
// requested network_enabled value, or the stream ends / the request times out.
func awaitOverviewNetworkEnabled(reader *bufio.Reader, want bool) error {
for {
event, data, err := readSSEEvent(reader)
if err != nil {
return err
}
if event != "overview" {
continue
}
var overview struct {
NetworkEnabled bool `json:"network_enabled"`
}
if err := json.Unmarshal(data, &overview); err != nil {
return err
}
if overview.NetworkEnabled == want {
return nil
}
}
}
// The overview SSE stream must reflect edits made after it opened. Before the
// fix it rebuilt every tick from the config snapshot captured when the stream
// opened, so toggling roaming data off was immediately overwritten by the stale
// "on" snapshot and the switch flapped. This test opens the stream with roaming
// data on, turns it off in the store, and requires the stream to keep reporting
// the new "off" state.
func TestHandleOverviewStreamReflectsConfigChanges(t *testing.T) {
previousInterval := overviewStreamInterval
overviewStreamInterval = 10 * time.Millisecond
t.Cleanup(func() { overviewStreamInterval = previousInterval })
ctx := context.Background()
database, err := store.Open(ctx, ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
if err := database.UpsertAppSetting(ctx, store.AppSetting{
Key: developer.EnabledSettingKey,
Value: []byte(`{"enabled":true}`),
}); err != nil {
t.Fatal(err)
}
if err := database.UpsertDevice(ctx, store.Device{ID: "dev1", Name: "Test device", NetworkEnabled: true}); err != nil {
t.Fatal(err)
}
server := &Server{store: database, logger: regionTestLogger(), developerEnabled: true}
mux := http.NewServeMux()
mux.HandleFunc("/stream", func(w http.ResponseWriter, r *http.Request) {
config, err := database.Device(r.Context(), "dev1")
if err != nil {
writeError(w, http.StatusNotFound, "not_found", err.Error())
return
}
server.handleOverviewStream(w, r, config, device.Device{}, false)
})
testServer := httptest.NewServer(mux)
t.Cleanup(testServer.Close)
requestCtx, cancel := context.WithTimeout(ctx, 10*time.Second)
t.Cleanup(cancel)
request, err := http.NewRequestWithContext(requestCtx, http.MethodGet, testServer.URL+"/stream", nil)
if err != nil {
t.Fatal(err)
}
response, err := http.DefaultClient.Do(request)
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = response.Body.Close() })
if response.StatusCode != http.StatusOK {
t.Fatalf("stream status = %d", response.StatusCode)
}
reader := bufio.NewReader(response.Body)
// The stream opens with roaming data enabled.
if err := awaitOverviewNetworkEnabled(reader, true); err != nil {
t.Fatalf("initial overview never reported network_enabled=true: %v", err)
}
// Turn roaming data off; the very next ticks must report the new state
// instead of replaying the stale enabled snapshot.
config, err := database.Device(ctx, "dev1")
if err != nil {
t.Fatal(err)
}
config.NetworkEnabled = false
if err := database.UpsertDevice(ctx, config); err != nil {
t.Fatal(err)
}
if err := awaitOverviewNetworkEnabled(reader, false); err != nil {
t.Fatalf("overview kept replaying stale network_enabled=true after the edit: %v", err)
}
}
// Turning roaming data off must be refused while an enabled export proxy is
// bound to the device; the user has to disable that binding first.
func TestHandleCellularDataRejectsDisableWhileExportProxyActive(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
if err := database.UpsertAppSetting(ctx, store.AppSetting{
Key: developer.EnabledSettingKey, Value: json.RawMessage(`{"enabled":true}`),
}); err != nil {
t.Fatal(err)
}
deviceConfig := store.Device{ID: "modem-1", Name: "modem-1", Interface: "wwan0", NetworkEnabled: true}
if err := database.UpsertDevice(ctx, deviceConfig); err != nil {
t.Fatal(err)
}
// Seed an already-enabled export proxy bound to the device. New only logs a
// warning when the Linux-only listener cannot start on this platform, so the
// enabled config still loads and the interlock sees it.
seeded, err := json.Marshal([]exportproxy.Config{{
ID: "proxy-1", Name: "proxy-1", DeviceID: "modem-1", Interface: "wwan0",
Mode: "socks5", ListenHost: "127.0.0.1", ListenPort: 1080, Enabled: true,
}})
if err != nil {
t.Fatal(err)
}
if err := database.UpsertAppSetting(ctx, store.AppSetting{Key: exportproxy.SettingKey, Value: seeded, Sensitive: true}); err != nil {
t.Fatal(err)
}
proxyManager, err := exportproxy.New(ctx, database, regionTestLogger(), "")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = proxyManager.Close() })
server := &Server{
store: database,
logger: regionTestLogger(),
developerEnabled: true,
exportProxy: proxyManager,
devices: fakeDeviceController{},
maxRequestBodyBytes: 1 << 20,
}
patchOff := func() *httptest.ResponseRecorder {
recorder := httptest.NewRecorder()
request := httptest.NewRequest(http.MethodPatch, "/api/devices/modem-1/cellular-data", strings.NewReader(`{"enabled":false}`))
request.Header.Set("Content-Type", "application/json")
if !server.handleCellularData(recorder, request, deviceConfig, "physical-1") {
t.Fatal("handleCellularData did not handle the request")
}
return recorder
}
// While the export proxy is enabled, turning roaming data off is rejected and
// the stored config keeps roaming data on.
recorder := patchOff()
if recorder.Code != http.StatusConflict {
t.Fatalf("disable with active proxy status = %d, body = %s", recorder.Code, recorder.Body)
}
var failure struct {
Error struct {
Code string `json:"code"`
} `json:"error"`
}
if err := json.Unmarshal(recorder.Body.Bytes(), &failure); err != nil {
t.Fatal(err)
}
if failure.Error.Code != "export_proxy_active" {
t.Fatalf("error code = %q, body = %s", failure.Error.Code, recorder.Body)
}
stored, err := database.Device(ctx, "modem-1")
if err != nil {
t.Fatal(err)
}
if !stored.NetworkEnabled {
t.Fatal("roaming data was turned off despite the active export proxy")
}
// Once the binding is disabled, the same request goes through.
proxies, err := proxyManager.Configs()
if err != nil || len(proxies) != 1 {
t.Fatalf("configs = %+v, %v", proxies, err)
}
disabled := proxies[0]
disabled.Enabled = false
if _, err := proxyManager.Update(ctx, disabled.ID, disabled); err != nil {
t.Fatal(err)
}
recorder = patchOff()
if recorder.Code != http.StatusOK {
t.Fatalf("disable after proxy off status = %d, body = %s", recorder.Code, recorder.Body)
}
stored, err = database.Device(ctx, "modem-1")
if err != nil {
t.Fatal(err)
}
if stored.NetworkEnabled {
t.Fatal("roaming data was not turned off after the export proxy was disabled")
}
}
+51
View File
@@ -0,0 +1,51 @@
package server
import (
"context"
"testing"
"time"
"vocat/internal/device"
"vocat/internal/store"
)
func TestConfiguredDeviceSummaryIgnoresVoWiFiRuntimeFromPreviousSIM(t *testing.T) {
database, err := store.Open(context.Background(), ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
if err := database.UpsertDevice(context.Background(), store.Device{ID: "ec20_1", Name: "EC20"}); err != nil {
t.Fatal(err)
}
if err := database.UpsertVoWiFiRuntime(context.Background(), store.VoWiFiRuntime{
DeviceID: "ec20_1",
Phase: "stopping",
ICCID: "89441000400128014257",
IMSI: "234159608751160",
TunnelReady: true,
IMSReady: true,
SMSReady: true,
LocalPhone: "+447386083638",
PhoneNumberSource: "ims_p_associated_uri",
UpdatedAt: time.Now().UTC(),
}); err != nil {
t.Fatal(err)
}
s := &Server{store: database}
entry := &device.Device{ID: "physical", Snapshot: &device.Snapshot{
ICCID: "89104100000028106378",
IMSI: "310380500712483",
}}
got := s.configuredDeviceSummary(store.Device{ID: "ec20_1"}, entry)
if got["vowifi_active"] != false {
t.Fatalf("vowifi_active = %#v", got["vowifi_active"])
}
if got["local_phone"] == "+447386083638" {
t.Fatalf("old phone leaked into current SIM summary: %#v", got)
}
runtime, ok := got["vowifi_runtime"].(map[string]any)
if !ok || runtime["phase"] != "idle" || runtime["iccid"] != "89104100000028106378" {
t.Fatalf("runtime = %#v", got["vowifi_runtime"])
}
}
+102
View File
@@ -0,0 +1,102 @@
package server
import (
"errors"
"net/http"
"strings"
"vocat/internal/exportproxy"
)
func (s *Server) routeExportProxyAPI(w http.ResponseWriter, r *http.Request, cleanPath string) bool {
if cleanPath != "export-proxies" && !strings.HasPrefix(cleanPath, "export-proxies/") {
return false
}
if !s.developerActive(r.Context()) || s.exportProxy == nil {
writeError(w, http.StatusForbidden, "developer_mode_required", "Export Proxy is available only in developer mode")
return true
}
segments := splitAPIPath(cleanPath)
if len(segments) == 1 {
switch r.Method {
case http.MethodGet:
configs, err := s.exportProxy.Configs()
if err != nil {
s.writeExportProxyError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{"configs": configs}})
case http.MethodPost:
var config exportproxy.Config
if err := s.decodeJSON(w, r, &config); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
created, err := s.exportProxy.Create(r.Context(), config)
if err != nil {
s.writeExportProxyError(w, err)
return true
}
writeJSON(w, http.StatusCreated, map[string]any{"data": created})
default:
w.Header().Set("Allow", "GET, POST")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
return true
}
if len(segments) == 2 && segments[1] == "status" {
if !requireMethod(w, r, http.MethodGet) {
return true
}
statuses, err := s.exportProxy.Status()
if err != nil {
s.writeExportProxyError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]any{"configs": statuses}})
return true
}
if len(segments) != 2 || strings.TrimSpace(segments[1]) == "" {
writeError(w, http.StatusNotFound, "not_found", "Export Proxy endpoint not found")
return true
}
id := segments[1]
switch r.Method {
case http.MethodPut:
var config exportproxy.Config
if err := s.decodeJSON(w, r, &config); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
updated, err := s.exportProxy.Update(r.Context(), id, config)
if err != nil {
s.writeExportProxyError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": updated})
case http.MethodDelete:
if err := s.exportProxy.Delete(r.Context(), id); err != nil {
s.writeExportProxyError(w, err)
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]bool{"deleted": true}})
default:
w.Header().Set("Allow", "PUT, DELETE")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
return true
}
func (s *Server) writeExportProxyError(w http.ResponseWriter, err error) {
switch {
case errors.Is(err, exportproxy.ErrDisabled):
writeError(w, http.StatusForbidden, "developer_mode_required", "Export Proxy is disabled")
case errors.Is(err, exportproxy.ErrNotFound):
writeError(w, http.StatusNotFound, "export_proxy_not_found", err.Error())
default:
writeError(w, http.StatusBadRequest, "export_proxy_invalid", err.Error())
}
}
+161
View File
@@ -0,0 +1,161 @@
package server
import (
"context"
"errors"
"io"
"net/http"
"os"
"strings"
"time"
)
const maxPluginUploadBytes int64 = 64 << 20
func (s *Server) routeExtensionAPI(w http.ResponseWriter, r *http.Request, cleanPath string) bool {
if cleanPath == "extensions" {
if s.extensions == nil {
writeError(w, http.StatusServiceUnavailable, "extensions_unavailable", "plugin manager is unavailable")
return true
}
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, map[string]any{"data": s.extensions.List()})
default:
w.Header().Set("Allow", "GET")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
return true
}
if cleanPath == "extensions/install-url" {
if s.extensions == nil {
writeError(w, http.StatusServiceUnavailable, "extensions_unavailable", "plugin manager is unavailable")
return true
}
if !requireMethod(w, r, http.MethodPost) {
return true
}
var request struct {
URL string `json:"url"`
SHA256 string `json:"sha256"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
ctx, cancel := contextWithTimeout(r, 60*time.Second)
defer cancel()
plugin, err := s.extensions.InstallURL(ctx, request.URL, request.SHA256)
if err != nil {
writeError(w, http.StatusBadRequest, "plugin_install_failed", err.Error())
return true
}
s.recordAudit(r.Context(), "admin", "plugin.install_url", "plugin", plugin.ID, "success", request.URL)
writeJSON(w, http.StatusCreated, map[string]any{"data": plugin})
return true
}
if cleanPath == "extensions/upload" {
if s.extensions == nil {
writeError(w, http.StatusServiceUnavailable, "extensions_unavailable", "plugin manager is unavailable")
return true
}
if !requireMethod(w, r, http.MethodPost) {
return true
}
r.Body = http.MaxBytesReader(w, r.Body, maxPluginUploadBytes+(1<<20))
if err := r.ParseMultipartForm(maxPluginUploadBytes); err != nil {
writeError(w, http.StatusBadRequest, "invalid_plugin_upload", "plugin upload must be multipart/form-data and no larger than 64 MiB")
return true
}
file, _, err := r.FormFile("package")
if err != nil {
writeError(w, http.StatusBadRequest, "invalid_plugin_upload", "multipart field package is required")
return true
}
defer file.Close()
plugin, err := s.extensions.Install(io.LimitReader(file, maxPluginUploadBytes+1), r.FormValue("sha256"))
if err != nil {
writeError(w, http.StatusBadRequest, "plugin_install_failed", err.Error())
return true
}
s.recordAudit(r.Context(), "admin", "plugin.upload", "plugin", plugin.ID, "success", "upload")
writeJSON(w, http.StatusCreated, map[string]any{"data": plugin})
return true
}
segments := splitAPIPath(cleanPath)
if len(segments) < 2 || segments[0] != "extensions" {
return false
}
if s.extensions == nil {
writeError(w, http.StatusServiceUnavailable, "extensions_unavailable", "plugin manager is unavailable")
return true
}
id := segments[1]
if len(segments) >= 3 && segments[2] == "backend" {
s.extensions.ProxyBackend(w, r, id)
return true
}
if len(segments) != 2 {
writeError(w, http.StatusNotFound, "not_found", "plugin endpoint not found")
return true
}
switch r.Method {
case http.MethodPut:
var request struct {
Enabled bool `json:"enabled"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return true
}
plugin, err := s.extensions.SetEnabled(id, request.Enabled)
if errors.Is(err, os.ErrNotExist) {
writeError(w, http.StatusNotFound, "plugin_not_found", "plugin not found")
return true
}
if err != nil {
writeError(w, http.StatusInternalServerError, "plugin_state_failed", err.Error())
return true
}
writeJSON(w, http.StatusOK, map[string]any{"data": plugin})
case http.MethodDelete:
if err := s.extensions.Uninstall(id); errors.Is(err, os.ErrNotExist) {
writeError(w, http.StatusNotFound, "plugin_not_found", "plugin not found")
} else if err != nil {
writeError(w, http.StatusInternalServerError, "plugin_uninstall_failed", err.Error())
} else {
writeJSON(w, http.StatusOK, map[string]any{"data": map[string]bool{"uninstalled": true}})
}
default:
w.Header().Set("Allow", "PUT, DELETE")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
return true
}
func (s *Server) handlePluginAsset(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodHead {
w.Header().Set("Allow", "GET, HEAD")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
return
}
if !s.requireAuthenticated(w, r) {
return
}
if s.extensions == nil {
http.NotFound(w, r)
return
}
path := strings.TrimPrefix(r.URL.Path, "/plugin-assets/")
parts := strings.SplitN(path, "/", 2)
if len(parts) != 2 || parts[0] == "" || parts[1] == "" {
http.NotFound(w, r)
return
}
s.extensions.ServeAsset(w, r, parts[0], parts[1])
}
func contextWithTimeout(r *http.Request, timeout time.Duration) (context.Context, context.CancelFunc) {
return context.WithTimeout(r.Context(), timeout)
}
+135 -8
View File
@@ -7,6 +7,7 @@ import (
"fmt"
"log/slog"
"net/http"
"os"
"runtime"
"strconv"
"strings"
@@ -14,13 +15,21 @@ import (
"vocat/internal/auth"
"vocat/internal/buildinfo"
"vocat/internal/developer"
"vocat/internal/i18n"
"vocat/internal/loghub"
"vocat/internal/store"
"vocat/internal/update"
)
func (s *Server) routeGeneralAPI(w http.ResponseWriter, r *http.Request) bool {
cleanPath := strings.Trim(strings.TrimPrefix(r.URL.Path, "/api"), "/")
if s.routeExtensionAPI(w, r, cleanPath) {
return true
}
if s.routeExportProxyAPI(w, r, cleanPath) {
return true
}
if s.routeSMSAPI(w, r, cleanPath) {
return true
}
@@ -45,6 +54,12 @@ func (s *Server) routeGeneralAPI(w http.ResponseWriter, r *http.Request) bool {
s.handlePasswordChange(w, r)
case "settings/preferences":
s.handleUIPreferences(w, r)
case "settings/https":
s.handleHTTPSSettings(w, r)
case "settings/https/certificate":
s.handleHTTPSCertificate(w, r)
case "settings/developer":
s.handleDeveloperSettings(w, r)
default:
return false
}
@@ -309,36 +324,148 @@ func (s *Server) handleSystemInfo(w http.ResponseWriter, r *http.Request) {
"os": runtime.GOOS,
"architecture": runtime.GOARCH,
"uptime": formatDuration(time.Since(s.startedAt)),
"developer": s.developerActive(r.Context()),
},
})
}
func (s *Server) developerActive(ctx context.Context) bool {
return s.developerEnabled && developer.Enabled(ctx, s.store)
}
func (s *Server) handleUpdateCheck(w http.ResponseWriter, r *http.Request) {
if !requireMethod(w, r, http.MethodGet) {
return
}
if strings.TrimSpace(s.updateRepository) == "" || s.updateCheck == nil {
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"available": false,
"version": buildinfo.Version,
"message": i18n.T("未配置受信任的软件更新源;不会从未知地址下载或执行文件。"),
},
})
return
}
ctx, cancel := context.WithTimeout(r.Context(), 15*time.Second)
defer cancel()
result, err := s.updateCheck(
ctx,
s.updateRepository,
s.updateToken,
buildinfo.Version,
)
if err != nil {
s.logger.Warn("check for updates failed", "repository", s.updateRepository, "error", err)
writeError(w, http.StatusBadGateway, "update_check_failed", err.Error())
return
}
message := ""
if result.Available {
message = result.ReleaseNotes
}
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"available": false,
"version": buildinfo.Version,
"message": i18n.T("未配置受信任的软件更新源;不会从未知地址下载或执行文件。"),
"available": result.Available,
"current_version": result.Current,
"version": result.Latest,
"message": message,
"repository": s.updateRepository,
"is_docker": runningInDocker(),
},
})
}
// handleUpdateApply deliberately performs no update. Without a configured,
// trusted update channel the product never downloads or executes code, so an
// apply request is acknowledged as a safe no-op rather than acted on.
func runningInDocker() bool {
if _, err := os.Stat("/.dockerenv"); err == nil {
return true
}
return strings.EqualFold(strings.TrimSpace(os.Getenv("VOCAT_CONTAINER")), "docker")
}
func (s *Server) handleUpdateApply(w http.ResponseWriter, r *http.Request) {
if !requireMethod(w, r, http.MethodPost) {
return
}
if strings.TrimSpace(s.updateRepository) == "" || s.updateApply == nil {
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"applied": false,
"message": i18n.T("未配置受信任的软件更新源;未执行任何更新。"),
},
})
return
}
if runningInDocker() {
writeError(w, http.StatusConflict, "container_update_required", "pull the latest container image and recreate the container")
return
}
s.updateMu.Lock()
if s.updateApplying {
s.updateMu.Unlock()
writeError(w, http.StatusConflict, "update_busy", "another update is already in progress")
return
}
s.updateApplying = true
s.updateMu.Unlock()
defer func() {
s.updateMu.Lock()
s.updateApplying = false
s.updateMu.Unlock()
}()
ctx, cancel := context.WithTimeout(r.Context(), 2*time.Minute)
defer cancel()
result, err := s.updateApply(ctx, s.logger, update.Options{
Repo: s.updateRepository,
Token: s.updateToken,
}, false)
if err != nil {
s.logger.Error("apply update failed", "repository", s.updateRepository, "error", err)
writeError(w, http.StatusBadGateway, "update_apply_failed", err.Error())
return
}
if !result.Applied {
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"applied": false,
"version": result.Latest,
"message": "The installed version is already current.",
},
})
return
}
// A binary update changes the trusted server code underneath every active
// browser/API session. Revoke every durable token before scheduling the
// restart and expire this client's cookies so all users must authenticate
// against the newly installed version.
if err := s.store.DeleteAllSessions(r.Context()); err != nil {
s.logger.Error("revoke sessions after update failed", "error", err)
writeError(w, http.StatusInternalServerError, "update_session_revocation_failed", "The update was installed, but active sessions could not be revoked; restart the service and sign in again.")
return
}
s.clearAuthCookies(w)
writeJSON(w, http.StatusOK, map[string]any{
"data": map[string]any{
"applied": false,
"message": i18n.T("未配置受信任的软件更新源;未执行任何更新。"),
"applied": true,
"version": result.Latest,
"reauthentication_required": true,
"message": "Update verified and installed; all sessions were revoked and the service is restarting.",
},
})
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
if s.updateRestart != nil {
restart := s.updateRestart
logger := s.logger
go func() {
time.Sleep(time.Second)
if err := restart(logger); err != nil {
logger.Error("restart after update failed", "error", err)
}
}()
}
}
func (s *Server) handlePasswordChange(w http.ResponseWriter, r *http.Request) {
+64
View File
@@ -0,0 +1,64 @@
package server
import (
"net/http"
"strconv"
)
func (s *Server) handleHTTPSSettings(w http.ResponseWriter, r *http.Request) {
if !s.developerEnabled {
writeError(w, http.StatusNotFound, "not_found", "resource not found")
return
}
if s.https == nil {
writeError(w, http.StatusServiceUnavailable, "https_unavailable", "self-signed HTTPS is unavailable")
return
}
switch r.Method {
case http.MethodGet:
writeJSON(w, http.StatusOK, map[string]any{"data": s.https.State(r.Host)})
case http.MethodPut:
var request struct {
Enabled bool `json:"enabled"`
}
if err := s.decodeJSON(w, r, &request); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request", err.Error())
return
}
state, err := s.https.SetEnabled(r.Context(), request.Enabled)
if err != nil {
writeError(w, http.StatusInternalServerError, "https_update_failed", err.Error())
return
}
state = s.https.State(r.Host)
s.recordAudit(r.Context(), "admin", "settings.https.update", "settings", "https", "success", map[bool]string{true: "enabled", false: "disabled"}[request.Enabled])
writeJSON(w, http.StatusOK, map[string]any{"data": state})
default:
w.Header().Set("Allow", "GET, PUT")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
}
}
func (s *Server) handleHTTPSCertificate(w http.ResponseWriter, r *http.Request) {
if !s.developerEnabled {
writeError(w, http.StatusNotFound, "not_found", "resource not found")
return
}
if !requireMethod(w, r, http.MethodGet) {
return
}
if s.https == nil {
writeError(w, http.StatusServiceUnavailable, "https_unavailable", "self-signed HTTPS is unavailable")
return
}
certificate, err := s.https.CertificatePEM()
if err != nil {
writeError(w, http.StatusInternalServerError, "certificate_unavailable", err.Error())
return
}
w.Header().Set("Content-Type", "application/x-pem-file")
w.Header().Set("Content-Disposition", `attachment; filename="vocat-selfsigned.crt"`)
w.Header().Set("Content-Length", strconv.Itoa(len(certificate)))
w.WriteHeader(http.StatusOK)
_, _ = w.Write(certificate)
}
+177
View File
@@ -0,0 +1,177 @@
package server
import (
"fmt"
"math"
"net"
"sync"
"time"
)
// liveNetWindow is how far back the "last minute" byte totals reach.
const liveNetWindow = time.Minute
// liveNetMaxGap bounds how far apart two samples may be before a rate computed
// across them stops being "live". The overview SSE ticks every two seconds, so
// a gap beyond this means the tab was closed or the device was idle; treat it
// as a fresh baseline instead of averaging a long dead interval.
const liveNetMaxGap = 15 * time.Second
// netIfSample is one cumulative counter reading for an interface.
type netIfSample struct {
at time.Time
rxCum uint64
txCum uint64
}
// liveNetDevice holds the per-device sampling state used to derive rates and
// trailing-window totals from cumulative interface counters.
type liveNetDevice struct {
prev netIfSample
hasPrev bool
window []netIfSample
}
// liveNetResult is one rendered snapshot of a device's live network state.
type liveNetResult struct {
ipv4 string
rxRate float64 // bytes/sec over the trailing sample interval
txRate float64
minuteRx int64 // bytes over the trailing liveNetWindow
minuteTx int64
status string // "", "waiting_sample", or "stale"
}
// liveNetTracker derives live rates and last-minute totals from cumulative
// /sys interface counters. It is driven on demand by the overview builders, so
// no separate goroutine is required; the SSE overview cadence keeps it warm.
type liveNetTracker struct {
mu sync.Mutex
devices map[string]*liveNetDevice
}
func newLiveNetTracker() *liveNetTracker {
return &liveNetTracker{devices: map[string]*liveNetDevice{}}
}
// sample reads the interface's current counters and addresses and returns the
// device's live network state. Interface addresses resolve even on the first
// call; rates and totals need a second reading, reported as waiting_sample.
func (t *liveNetTracker) sample(deviceID, iface string, now time.Time) liveNetResult {
ipv4 := netIfAddrs(iface)
rxCum, txCum, err := netIfCounters(iface)
if err != nil {
// The interface briefly disappears while QMI reconnects. Drop the
// baseline so the next good read starts fresh rather than counting the
// reconnect as one giant delta.
t.mu.Lock()
delete(t.devices, deviceID)
t.mu.Unlock()
return liveNetResult{ipv4: ipv4, status: "stale"}
}
rxRate, txRate, minuteRx, minuteTx, status := t.record(deviceID, rxCum, txCum, now)
return liveNetResult{
ipv4: ipv4,
rxRate: rxRate, txRate: txRate,
minuteRx: minuteRx, minuteTx: minuteTx,
status: status,
}
}
// record folds one cumulative counter reading into the device's sampling state
// and returns the derived rates and trailing-window totals. It is pure (no
// interface I/O) so the rate/window logic is unit-testable.
func (t *liveNetTracker) record(deviceID string, rxCum, txCum uint64, now time.Time) (rxRate, txRate float64, minuteRx, minuteTx int64, status string) {
t.mu.Lock()
defer t.mu.Unlock()
d := t.devices[deviceID]
if d == nil {
d = &liveNetDevice{}
t.devices[deviceID] = d
}
current := netIfSample{at: now, rxCum: rxCum, txCum: txCum}
// First sighting, a counter reset (interface reconnected), or a gap too
// long to average honestly: establish a baseline and wait for the next
// reading before reporting a rate.
if !d.hasPrev || rxCum < d.prev.rxCum || txCum < d.prev.txCum || now.Sub(d.prev.at) > liveNetMaxGap {
d.prev = current
d.hasPrev = true
d.window = []netIfSample{current}
return 0, 0, 0, 0, "waiting_sample"
}
if elapsed := now.Sub(d.prev.at).Seconds(); elapsed > 0 {
rxRate = float64(rxCum-d.prev.rxCum) / elapsed
txRate = float64(txCum-d.prev.txCum) / elapsed
}
d.prev = current
d.window = append(d.window, current)
// Drop samples outside the trailing window, then measure totals against
// the oldest surviving reading.
cutoff := now.Add(-liveNetWindow)
kept := d.window[:0]
for _, s := range d.window {
if !s.at.Before(cutoff) {
kept = append(kept, s)
}
}
d.window = kept
minuteRx = int64(rxCum - d.window[0].rxCum)
minuteTx = int64(txCum - d.window[0].txCum)
return rxRate, txRate, minuteRx, minuteTx, ""
}
// netIfAddrs returns the interface's first global IPv4 address. It uses only
// the net package, so it compiles on every platform; on hosts without the
// interface it returns an empty string.
func netIfAddrs(iface string) (ipv4 string) {
if iface == "" {
return ""
}
netIf, err := net.InterfaceByName(iface)
if err != nil {
return ""
}
addrs, err := netIf.Addrs()
if err != nil {
return ""
}
for _, addr := range addrs {
var ip net.IP
switch a := addr.(type) {
case *net.IPNet:
ip = a.IP
case *net.IPAddr:
ip = a.IP
}
if ip == nil || ip.IsLoopback() {
continue
}
if v4 := ip.To4(); v4 != nil {
return v4.String()
}
}
return ""
}
// formatLiveBytes mirrors the SPA's formatBytes so the live strings match the
// chart's formatting: 1024-based units, rounded once the value reaches 100.
func formatLiveBytes(value float64) string {
if math.IsNaN(value) || math.IsInf(value, 0) || value < 0 {
value = 0
}
units := []string{"B", "KB", "MB", "GB", "TB"}
size := value
unit := 0
for size >= 1024 && unit < len(units)-1 {
size /= 1024
unit++
}
if size >= 100 {
return fmt.Sprintf("%.0f %s", size, units[unit])
}
return fmt.Sprintf("%.1f %s", size, units[unit])
}
+143
View File
@@ -0,0 +1,143 @@
package server
import (
"testing"
"time"
)
// record is the pure rate/window core of the tracker; these tests drive it
// directly with synthetic cumulative counters, no interface I/O involved.
func TestLiveNetRecordFirstSampleWaits(t *testing.T) {
tracker := newLiveNetTracker()
now := time.Unix(1_700_000_000, 0)
_, _, _, _, status := tracker.record("dev1", 1000, 500, now)
if status != "waiting_sample" {
t.Fatalf("first sample status = %q, want waiting_sample", status)
}
}
func TestLiveNetRecordComputesRateAndMinute(t *testing.T) {
tracker := newLiveNetTracker()
base := time.Unix(1_700_000_000, 0)
tracker.record("dev1", 1000, 500, base)
rxRate, txRate, minuteRx, minuteTx, status := tracker.record("dev1", 2000, 700, base.Add(2*time.Second))
if status != "" {
t.Fatalf("second sample status = %q, want empty", status)
}
// 1000 rx bytes and 200 tx bytes over 2s.
if rxRate != 500 {
t.Errorf("rxRate = %v, want 500", rxRate)
}
if txRate != 100 {
t.Errorf("txRate = %v, want 100", txRate)
}
if minuteRx != 1000 {
t.Errorf("minuteRx = %v, want 1000", minuteRx)
}
if minuteTx != 200 {
t.Errorf("minuteTx = %v, want 200", minuteTx)
}
}
func TestLiveNetRecordUsesActualElapsed(t *testing.T) {
tracker := newLiveNetTracker()
base := time.Unix(1_700_000_000, 0)
tracker.record("dev1", 0, 0, base)
// A 4s gap (not the usual 2s tick) must divide by 4, not 2.
rxRate, _, _, _, status := tracker.record("dev1", 400, 0, base.Add(4*time.Second))
if status != "" {
t.Fatalf("status = %q, want empty", status)
}
if rxRate != 100 {
t.Errorf("rxRate = %v, want 100", rxRate)
}
}
func TestLiveNetRecordCounterResetRebaselines(t *testing.T) {
tracker := newLiveNetTracker()
base := time.Unix(1_700_000_000, 0)
tracker.record("dev1", 5000, 5000, base)
tracker.record("dev1", 6000, 6000, base.Add(2*time.Second))
// Counter drops (interface reconnected): must re-baseline, not go negative.
_, _, _, _, status := tracker.record("dev1", 100, 100, base.Add(4*time.Second))
if status != "waiting_sample" {
t.Fatalf("after reset status = %q, want waiting_sample", status)
}
}
func TestLiveNetRecordLongGapRebaselines(t *testing.T) {
tracker := newLiveNetTracker()
base := time.Unix(1_700_000_000, 0)
tracker.record("dev1", 1000, 1000, base)
// Gap beyond liveNetMaxGap (tab closed / idle): treat as fresh baseline.
_, _, _, _, status := tracker.record("dev1", 2000, 2000, base.Add(liveNetMaxGap+time.Second))
if status != "waiting_sample" {
t.Fatalf("after long gap status = %q, want waiting_sample", status)
}
}
func TestLiveNetRecordSlidesWindow(t *testing.T) {
tracker := newLiveNetTracker()
base := time.Unix(1_700_000_000, 0)
// One sample every 2s, rx climbing 100 bytes each tick (50 B/s).
tracker.record("dev1", 0, 0, base)
var minuteRx int64
var status string
for i := 1; i <= 31; i++ {
now := base.Add(time.Duration(2*i) * time.Second) // t=2s .. t=62s
_, _, minuteRx, _, status = tracker.record("dev1", uint64(100*i), 0, now)
}
if status != "" {
t.Fatalf("status = %q, want empty", status)
}
// At t=62s the cutoff is t=2s, so the t=0 baseline has slid out. The window
// now spans t=2s..t=62s = 60s and 30 ticks of 100 bytes.
if minuteRx != 3000 {
t.Errorf("minuteRx = %v, want 3000 (only trailing window)", minuteRx)
}
}
func TestLiveNetRecordTracksDevicesIndependently(t *testing.T) {
tracker := newLiveNetTracker()
base := time.Unix(1_700_000_000, 0)
tracker.record("a", 1000, 0, base)
tracker.record("b", 9000, 0, base)
rxRateA, _, _, _, _ := tracker.record("a", 2000, 0, base.Add(2*time.Second))
rxRateB, _, _, _, _ := tracker.record("b", 9100, 0, base.Add(2*time.Second))
if rxRateA != 500 {
t.Errorf("device a rxRate = %v, want 500", rxRateA)
}
if rxRateB != 50 {
t.Errorf("device b rxRate = %v, want 50", rxRateB)
}
}
func TestFormatLiveBytes(t *testing.T) {
cases := []struct {
in float64
want string
}{
{0, "0.0 B"},
{512, "512 B"},
{1023, "1023 B"},
{1024, "1.0 KB"},
{1536, "1.5 KB"},
{100 * 1024, "100 KB"},
{5 * 1024 * 1024, "5.0 MB"},
{3 * 1024 * 1024 * 1024, "3.0 GB"},
{-5, "0.0 B"},
}
for _, c := range cases {
if got := formatLiveBytes(c.in); got != c.want {
t.Errorf("formatLiveBytes(%v) = %q, want %q", c.in, got, c.want)
}
}
}
+40
View File
@@ -0,0 +1,40 @@
//go:build linux
package server
import (
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
)
// netIfCounters reads the interface's cumulative rx/tx byte counters from
// /sys/class/net. The interface briefly disappears while QMI reconnects, in
// which case an error is returned and the caller re-baselines.
func netIfCounters(iface string) (uint64, uint64, error) {
if strings.TrimSpace(iface) == "" {
return 0, 0, fmt.Errorf("interface name is empty")
}
read := func(counter string) (uint64, error) {
raw, err := os.ReadFile(filepath.Join("/sys/class/net", iface, "statistics", counter))
if err != nil {
return 0, err
}
parsed, err := strconv.ParseUint(strings.TrimSpace(string(raw)), 10, 64)
if err != nil {
return 0, fmt.Errorf("parse %s %s counter: %w", iface, counter, err)
}
return parsed, nil
}
rxBytes, err := read("rx_bytes")
if err != nil {
return 0, 0, err
}
txBytes, err := read("tx_bytes")
if err != nil {
return 0, 0, err
}
return rxBytes, txBytes, nil
}
+11
View File
@@ -0,0 +1,11 @@
//go:build !linux
package server
import "fmt"
// netIfCounters is only meaningful on the Linux deployment target; elsewhere
// there is no cellular /sys interface to read.
func netIfCounters(string) (uint64, uint64, error) {
return 0, 0, fmt.Errorf("interface counters are only available on Linux")
}
+96
View File
@@ -0,0 +1,96 @@
package server
import (
"context"
"net/http"
"strings"
"time"
"vocat/internal/exportproxy"
"vocat/internal/store"
)
type cachedPublicIP struct {
ICCID string
Info exportproxy.PublicIPInfo
}
type publicIPResponse struct {
Detected bool `json:"detected"`
exportproxy.PublicIPInfo
}
func (s *Server) clearPublicIP(deviceID string) {
s.publicIPMu.Lock()
delete(s.publicIPs, strings.TrimSpace(deviceID))
s.publicIPMu.Unlock()
}
func (s *Server) loadPublicIP(deviceID, iccid string) (exportproxy.PublicIPInfo, bool) {
deviceID = strings.TrimSpace(deviceID)
iccid = strings.TrimSpace(iccid)
s.publicIPMu.RLock()
entry, ok := s.publicIPs[deviceID]
s.publicIPMu.RUnlock()
if !ok {
return exportproxy.PublicIPInfo{}, false
}
// A missing live ICCID means the modem is resetting or no card is present.
// A different ICCID means the SIM/eSIM profile changed. Either transition
// invalidates the old cellular exit immediately.
if iccid == "" || !strings.EqualFold(strings.TrimSpace(entry.ICCID), iccid) {
s.clearPublicIP(deviceID)
return exportproxy.PublicIPInfo{}, false
}
return entry.Info, true
}
func (s *Server) savePublicIP(deviceID, iccid string, info exportproxy.PublicIPInfo) {
s.publicIPMu.Lock()
s.publicIPs[strings.TrimSpace(deviceID)] = cachedPublicIP{
ICCID: strings.TrimSpace(iccid),
Info: info,
}
s.publicIPMu.Unlock()
}
func (s *Server) handleCellularPublicIP(w http.ResponseWriter, r *http.Request, config store.Device, iccid string) bool {
if r.Method != http.MethodGet && r.Method != http.MethodPost {
w.Header().Set("Allow", "GET, POST")
writeError(w, http.StatusMethodNotAllowed, "method_not_allowed", "method not allowed")
return true
}
if !s.developerActive(r.Context()) {
writeError(w, http.StatusForbidden, "developer_mode_required", "public IP detection through roaming data is available only in developer mode")
return true
}
w.Header().Set("Cache-Control", "no-store")
if r.Method == http.MethodGet {
info, ok := s.loadPublicIP(config.ID, iccid)
writeJSON(w, http.StatusOK, map[string]any{"data": publicIPResponse{Detected: ok, PublicIPInfo: info}})
return true
}
if !config.NetworkEnabled {
writeError(w, http.StatusConflict, "cellular_data_disabled", "enable roaming data before detecting its public IP")
return true
}
if strings.TrimSpace(iccid) == "" {
writeError(w, http.StatusConflict, "sim_identity_unavailable", "the modem has no current ICCID; refresh it before detecting the public IP")
return true
}
if strings.TrimSpace(config.Interface) == "" {
writeError(w, http.StatusConflict, "cellular_interface_missing", "the device has no cellular network interface")
return true
}
ctx, cancel := context.WithTimeout(r.Context(), 15*time.Second)
defer cancel()
info, err := exportproxy.LookupPublicIP(ctx, config.Interface)
if err != nil {
s.logger.Warn("detect roaming public IP failed", "device_id", config.ID, "interface", config.Interface, "error", err)
writeError(w, http.StatusBadGateway, "public_ip_lookup_failed", err.Error())
return true
}
s.savePublicIP(config.ID, iccid, info)
writeJSON(w, http.StatusOK, map[string]any{"data": publicIPResponse{Detected: true, PublicIPInfo: info}})
return true
}
+33
View File
@@ -0,0 +1,33 @@
package server
import (
"testing"
"vocat/internal/exportproxy"
)
func TestPublicIPCacheFollowsCurrentICCID(t *testing.T) {
server := &Server{publicIPs: make(map[string]cachedPublicIP)}
want := exportproxy.PublicIPInfo{IP: "203.0.113.8", CountryCode: "GB"}
server.savePublicIP("ec20", "8944100001", want)
got, ok := server.loadPublicIP("ec20", "8944100001")
if !ok || got != want {
t.Fatalf("loadPublicIP() = (%+v, %v), want (%+v, true)", got, ok, want)
}
if _, ok := server.loadPublicIP("ec20", "8944100002"); ok {
t.Fatal("cache survived an ICCID change")
}
if _, ok := server.loadPublicIP("ec20", "8944100001"); ok {
t.Fatal("stale cache was not deleted after an ICCID change")
}
}
func TestPublicIPCacheClearsWhileModemIsResetting(t *testing.T) {
server := &Server{publicIPs: make(map[string]cachedPublicIP)}
server.savePublicIP("ec20", "8944100001", exportproxy.PublicIPInfo{IP: "203.0.113.8", CountryCode: "GB"})
if _, ok := server.loadPublicIP("ec20", ""); ok {
t.Fatal("cache survived a missing live ICCID")
}
}
+8 -2
View File
@@ -21,6 +21,8 @@ import (
// USSD, and USB-net results are configurable for the feature endpoint tests.
type fakeDeviceController struct {
entry device.Device
atResponse modem.Response
atErr error
scanResult device.OperatorScanResult
scanErr error
ussdResult device.USSDResult
@@ -43,11 +45,11 @@ func (f fakeDeviceController) Refresh(context.Context, string) (device.Snapshot,
return device.Snapshot{}, nil
}
func (f fakeDeviceController) ExecuteAT(context.Context, string, string) (modem.Response, error) {
return modem.Response{}, nil
return f.atResponse, f.atErr
}
func (f fakeDeviceController) Reboot(context.Context, string) error { return nil }
func (f fakeDeviceController) USSD(context.Context, string, string) (device.USSDResult, error) {
return device.USSDResult{}, nil
return f.ussdResult, f.ussdErr
}
func (f fakeDeviceController) ContinueUSSD(context.Context, string, string) (device.USSDResult, error) {
return f.ussdResult, f.ussdErr
@@ -74,6 +76,10 @@ func (f fakeDeviceController) OperatorSelection(context.Context, string) (device
func (f fakeDeviceController) SetOperatorSelection(context.Context, string, bool, string, *int) (device.OperatorSelection, error) {
return device.OperatorSelection{}, nil
}
func (f fakeDeviceController) ReRegisterOperator(context.Context, string) (device.OperatorSelection, error) {
return device.OperatorSelection{}, nil
}
func (f fakeDeviceController) ScanOperators(context.Context, string) (device.OperatorScanResult, error) {
return f.scanResult, f.scanErr
}
+73 -12
View File
@@ -18,8 +18,12 @@ import (
"time"
"vocat/internal/auth"
"vocat/internal/exportproxy"
"vocat/internal/extensions"
"vocat/internal/httpsmode"
"vocat/internal/loghub"
"vocat/internal/store"
"vocat/internal/update"
"vocat/internal/vowifi"
)
@@ -39,6 +43,12 @@ type Options struct {
Logger *slog.Logger
SecureCookies bool
MaxRequestBodyBytes int64
Extensions *extensions.Manager
ExportProxy *exportproxy.Manager
DeveloperEnabled bool
UpdateRepository string
UpdateToken string
HTTPS *httpsmode.Manager
}
// Server is the single HTTP handler for the JSON API and embedded SPA.
@@ -60,6 +70,20 @@ type Server struct {
accessMu sync.RWMutex
access parsedAccessConfig
loginLimiter *loginRateLimiter
extensions *extensions.Manager
exportProxy *exportproxy.Manager
developerEnabled bool
updateRepository string
updateToken string
updateCheck func(context.Context, string, string, string) (update.CheckResult, error)
updateApply func(context.Context, *slog.Logger, update.Options, bool) (update.CheckResult, error)
updateRestart func(*slog.Logger) error
updateMu sync.Mutex
updateApplying bool
https *httpsmode.Manager
netTraffic *liveNetTracker
publicIPMu sync.RWMutex
publicIPs map[string]cachedPublicIP
}
func New(options Options) (*Server, error) {
@@ -82,6 +106,9 @@ func New(options Options) (*Server, error) {
if options.MaxRequestBodyBytes <= 0 {
options.MaxRequestBodyBytes = 1 << 20
}
if strings.TrimSpace(options.UpdateRepository) == "" {
options.UpdateRepository = update.DefaultRepository
}
server := &Server{
store: options.Store,
@@ -98,6 +125,17 @@ func New(options Options) (*Server, error) {
startedAt: time.Now().UTC(),
websheets: newWebsheetManager(),
loginLimiter: newLoginRateLimiter(),
extensions: options.Extensions,
exportProxy: options.ExportProxy,
developerEnabled: options.DeveloperEnabled,
updateRepository: strings.TrimSpace(options.UpdateRepository),
updateToken: strings.TrimSpace(options.UpdateToken),
https: options.HTTPS,
netTraffic: newLiveNetTracker(),
publicIPs: make(map[string]cachedPublicIP),
updateCheck: update.CheckLatest,
updateApply: update.ApplyLatest,
updateRestart: update.RestartService,
}
server.loadAccessConfig(context.Background())
server.loadUILanguage(context.Background())
@@ -110,6 +148,7 @@ func New(options Options) (*Server, error) {
mux.HandleFunc("/api", server.handleAPI)
mux.HandleFunc("/api/", server.handleAPI)
mux.HandleFunc("/websheets/", server.handleWebsheet)
mux.HandleFunc("/plugin-assets/", server.handlePluginAsset)
mux.HandleFunc("/", server.handleSPA)
server.handler = server.recoverPanics(
@@ -127,6 +166,17 @@ type VoWiFiController interface {
RequestReconnect(string) (vowifi.State, error)
}
type VoWiFiCallController interface {
Calls(string) ([]vowifi.Call, error)
DialCall(context.Context, string, string) (vowifi.Call, error)
AnswerCall(context.Context, string, string) (vowifi.Call, error)
HangupCall(context.Context, string, string) error
}
type VoWiFiCallMediaController interface {
CallMedia(context.Context, string, string) (vowifi.CallMedia, error)
}
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
s.handler.ServeHTTP(w, r)
}
@@ -227,8 +277,7 @@ func (s *Server) handleSession(w http.ResponseWriter, r *http.Request) {
}
session, csrfToken, err := s.auth.CSRFToken(r.Context(), sessionToken, existingCSRF)
if errors.Is(err, auth.ErrUnauthorized) {
s.clearAuthCookies(w)
writeError(w, http.StatusUnauthorized, "unauthorized", "authentication is required")
s.authenticationRequired(w, r)
return
}
if err != nil {
@@ -263,8 +312,7 @@ func (s *Server) handleLogout(w http.ResponseWriter, r *http.Request) {
if _, err := s.auth.ValidateCSRF(r.Context(), sessionToken, csrfToken); err != nil {
switch {
case errors.Is(err, auth.ErrUnauthorized):
s.clearAuthCookies(w)
writeError(w, http.StatusUnauthorized, "unauthorized", "authentication is required")
s.authenticationRequired(w, r)
case errors.Is(err, auth.ErrInvalidCSRF):
writeError(w, http.StatusForbidden, "invalid_csrf", "CSRF validation failed")
default:
@@ -313,8 +361,7 @@ func (s *Server) handleAPI(w http.ResponseWriter, r *http.Request) {
if _, err := s.auth.ValidateCSRF(r.Context(), sessionToken, csrfToken); err != nil {
switch {
case errors.Is(err, auth.ErrUnauthorized):
s.clearAuthCookies(w)
writeError(w, http.StatusUnauthorized, "unauthorized", "authentication is required")
s.authenticationRequired(w, r)
case errors.Is(err, auth.ErrInvalidCSRF):
writeError(w, http.StatusForbidden, "invalid_csrf", "CSRF validation failed")
default:
@@ -386,7 +433,7 @@ func (s *Server) decodeJSON(w http.ResponseWriter, r *http.Request, destination
func (s *Server) sessionToken(w http.ResponseWriter, r *http.Request) (string, bool) {
cookie, err := r.Cookie(sessionCookieName)
if err != nil || cookie.Value == "" {
writeError(w, http.StatusUnauthorized, "unauthorized", "authentication is required")
s.authenticationRequired(w, r)
return "", false
}
return cookie.Value, true
@@ -399,8 +446,7 @@ func (s *Server) requireAuthenticated(w http.ResponseWriter, r *http.Request) bo
}
if _, err := s.auth.Authenticate(r.Context(), sessionToken); err != nil {
if errors.Is(err, auth.ErrUnauthorized) {
s.clearAuthCookies(w)
writeError(w, http.StatusUnauthorized, "unauthorized", "authentication is required")
s.authenticationRequired(w, r)
} else {
s.logger.Error("request authentication failed", "error", err)
writeError(w, http.StatusInternalServerError, "internal_error", "an internal error occurred")
@@ -410,6 +456,21 @@ func (s *Server) requireAuthenticated(w http.ResponseWriter, r *http.Request) bo
return true
}
// authenticationRequired preserves JSON semantics for API clients while
// making a direct browser navigation land on the login screen instead of a
// raw {"error":...} document. Frontend fetches explicitly request JSON and
// are handled by the shared vocat:unauthorized event.
func (s *Server) authenticationRequired(w http.ResponseWriter, r *http.Request) {
s.clearAuthCookies(w)
w.Header().Set("Cache-Control", "no-store")
if (r.Method == http.MethodGet || r.Method == http.MethodHead) &&
strings.Contains(strings.ToLower(r.Header.Get("Accept")), "text/html") {
http.Redirect(w, r, "/login", http.StatusSeeOther)
return
}
writeError(w, http.StatusUnauthorized, "unauthorized", "authentication is required")
}
func (s *Server) validateDoubleSubmitCSRF(w http.ResponseWriter, r *http.Request) (string, bool) {
headerToken := r.Header.Get(csrfHeaderName)
cookie, err := r.Cookie(csrfCookieName)
@@ -527,8 +588,8 @@ func (s *Server) securityHeaders(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "same-origin")
w.Header().Set("Permissions-Policy", "camera=(), microphone=(), geolocation=()")
if strings.HasPrefix(r.URL.Path, "/websheets/") {
w.Header().Set("Permissions-Policy", "camera=(), microphone=(self), geolocation=()")
if strings.HasPrefix(r.URL.Path, "/websheets/") || strings.HasPrefix(r.URL.Path, "/plugin-assets/") {
// The self-hosted E911 websheet is embedded in an iframe by the SPA, so
// it must be frameable same-origin. Every other route stays DENY.
w.Header().Set("X-Frame-Options", "SAMEORIGIN")
@@ -549,7 +610,7 @@ func (s *Server) securityHeaders(next http.Handler) http.Handler {
"img-src 'self' data:; connect-src 'self'",
)
}
if s.secureCookies {
if s.secureCookies && s.https == nil {
w.Header().Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
}
next.ServeHTTP(w, r)
+24
View File
@@ -241,6 +241,30 @@ func TestUnifiedAPIErrors(t *testing.T) {
}
}
func TestUnauthenticatedBrowserNavigationRedirectsToLogin(t *testing.T) {
app := newTestApplication(t)
client := *app.client
client.CheckRedirect = func(_ *http.Request, _ []*http.Request) error {
return http.ErrUseLastResponse
}
request, err := http.NewRequest(http.MethodGet, app.server.URL+"/api/devices", nil)
if err != nil {
t.Fatal(err)
}
request.Header.Set("Accept", "text/html,application/xhtml+xml")
response, err := client.Do(request)
if err != nil {
t.Fatal(err)
}
defer response.Body.Close()
if response.StatusCode != http.StatusSeeOther {
t.Fatalf("navigation status = %d", response.StatusCode)
}
if location := response.Header.Get("Location"); location != "/login" {
t.Fatalf("navigation location = %q", location)
}
}
func TestNewRequiresIndex(t *testing.T) {
database, err := store.Open(context.Background(), ":memory:")
if err != nil {
+81 -25
View File
@@ -21,6 +21,7 @@ import (
"sort"
"strconv"
"strings"
"sync"
"time"
"vocat/internal/store"
@@ -255,8 +256,8 @@ func validateNotificationField(
return fmt.Errorf("%s is too long or contains invalid characters", field)
}
if name == "base_url" && value != "" {
if _, err := parseOutboundURL(value, true); err != nil {
return fmt.Errorf("%s must be an absolute HTTPS URL", field)
if _, err := telegramAPIURL(value, "123456:validation-token", "sendMessage"); err != nil {
return fmt.Errorf("%s must be an absolute HTTPS URL or a URL template with two %%s placeholders", field)
}
}
if name == "proxy" && value != "" {
@@ -551,8 +552,8 @@ func validateNotificationTestConfig(channel string, config map[string]any) error
return errors.New("telegram.chat_id is required")
}
if baseURL := configString(config, "base_url"); baseURL != "" {
if _, err := parseOutboundURL(baseURL, true); err != nil {
return errors.New("telegram.base_url must be an absolute HTTPS URL")
if _, err := telegramAPIURL(baseURL, token, "sendMessage"); err != nil {
return errors.New("telegram.base_url must be an absolute HTTPS URL or a URL template with two %s placeholders")
}
}
case "email":
@@ -660,19 +661,11 @@ func sendBarkNotificationTest(ctx context.Context, config map[string]any) error
}
func sendTelegramNotificationTest(ctx context.Context, config map[string]any) error {
baseURL := configString(config, "base_url")
if baseURL == "" {
baseURL = "https://api.telegram.org"
}
parsed, err := validateOutboundURL(ctx, baseURL, true)
token := configString(config, "bot_token")
parsed, err := validateTelegramAPIURL(ctx, configString(config, "base_url"), token, "sendMessage")
if err != nil {
return err
}
token := configString(config, "bot_token")
parsed.Path = strings.TrimRight(parsed.Path, "/") + "/bot" + token + "/sendMessage"
parsed.RawPath = ""
parsed.RawQuery = ""
parsed.Fragment = ""
client, err := restrictedHTTPClient(ctx, 6*time.Second, configString(config, "proxy"))
if err != nil {
return err
@@ -940,18 +933,77 @@ func dialRestricted(
if err != nil {
return nil, err
}
dialer := net.Dialer{Timeout: clampNotificationTimeout(timeout)}
var failures []error
for _, ip := range addresses {
connection, err := dialer.DialContext(
ctx,
network,
net.JoinHostPort(ip.String(), port),
)
if err == nil {
return connection, nil
perAddress := clampNotificationTimeout(timeout)
stagger := 300 * time.Millisecond
if perAddress < stagger {
stagger = perAddress / 2
}
raceContext, cancel := context.WithCancel(ctx)
defer cancel()
type attempt struct {
conn net.Conn
err error
}
resultCh := make(chan attempt, len(addresses))
var wg sync.WaitGroup
launcher := time.NewTicker(stagger)
defer launcher.Stop()
for index, ip := range addresses {
if index > 0 {
select {
case <-raceContext.Done():
break
case <-launcher.C:
}
}
failures = append(failures, err)
if raceContext.Err() != nil {
break
}
ip := ip
wg.Add(1)
go func() {
defer wg.Done()
dialer := net.Dialer{Timeout: perAddress}
connection, dialErr := dialer.DialContext(
raceContext,
network,
net.JoinHostPort(ip.String(), port),
)
if dialErr != nil {
resultCh <- attempt{err: dialErr}
return
}
if raceContext.Err() != nil {
connection.Close()
resultCh <- attempt{err: raceContext.Err()}
return
}
resultCh <- attempt{conn: connection}
}()
}
go func() {
wg.Wait()
close(resultCh)
}()
var failures []error
for result := range resultCh {
if result.conn != nil {
cancel()
return result.conn, nil
}
if result.err != nil && !errors.Is(result.err, context.Canceled) {
failures = append(failures, result.err)
}
if ctx.Err() != nil {
return nil, ctx.Err()
}
}
if len(failures) == 0 {
return nil, ctx.Err()
}
return nil, fmt.Errorf("dial public notification destination: %w", errors.Join(failures...))
}
@@ -1267,6 +1319,10 @@ func (s *Server) handleTrafficAnalysis(w http.ResponseWriter, r *http.Request) {
if !requireMethod(w, r, http.MethodGet) {
return
}
if !s.developerActive(r.Context()) {
writeError(w, http.StatusForbidden, "developer_mode_required", "traffic analysis is available only in developer mode")
return
}
rangeName := strings.ToLower(strings.TrimSpace(r.URL.Query().Get("range")))
if rangeName == "" {
rangeName = "day"
+39
View File
@@ -14,6 +14,7 @@ import (
"testing"
"time"
"vocat/internal/developer"
"vocat/internal/store"
)
@@ -211,6 +212,30 @@ func TestNotificationSettingsRejectsUnknownAndMalformedInput(t *testing.T) {
}
}
func TestNotificationSettingsAcceptsTelegramReverseProxyTemplate(t *testing.T) {
test := newSettingsAPITest(t)
recorder := test.request(
t,
http.MethodPut,
"/api/settings/notifications",
`{"telegram":{"enabled":true,"bot_token":"123456:abcdefghijklmnopqrstuvwxyz","chat_id":"1","base_url":"https://telegram.example.com/bot%s/%s"}}`,
)
if recorder.Code != http.StatusOK {
t.Fatalf("PUT status = %d, body = %s", recorder.Code, recorder.Body)
}
stored, err := test.database.NotificationSetting(context.Background(), "telegram")
if err != nil {
t.Fatal(err)
}
var config map[string]any
if err := json.Unmarshal(stored.Config, &config); err != nil {
t.Fatal(err)
}
if config["base_url"] != "https://telegram.example.com/bot%s/%s" {
t.Fatalf("stored Telegram base URL = %#v", config["base_url"])
}
}
func TestNotificationTestsBlockSSRFAndUnsupportedChannels(t *testing.T) {
test := newSettingsAPITest(t)
var webhookHits atomic.Int32
@@ -441,6 +466,12 @@ func TestCardPolicyDefaultValidationAndPersistence(t *testing.T) {
func TestTrafficAnalysisUsesAndAggregatesStoredBuckets(t *testing.T) {
test := newSettingsAPITest(t)
test.server.developerEnabled = true
if err := test.database.UpsertAppSetting(context.Background(), store.AppSetting{
Key: developer.EnabledSettingKey, Value: json.RawMessage(`{"enabled":true}`),
}); err != nil {
t.Fatal(err)
}
period := time.Now().UTC().Add(-time.Hour).Truncate(time.Minute)
for _, bucket := range []store.TrafficBucket{
{
@@ -493,6 +524,14 @@ func TestTrafficAnalysisUsesAndAggregatesStoredBuckets(t *testing.T) {
}
}
func TestTrafficAnalysisIsUnavailableOutsideDeveloperMode(t *testing.T) {
test := newSettingsAPITest(t)
recorder := test.request(t, http.MethodGet, "/api/traffic/analysis?range=week", "")
if recorder.Code != http.StatusForbidden {
t.Fatalf("traffic status = %d, want %d; body = %s", recorder.Code, http.StatusForbidden, recorder.Body)
}
}
func TestNotificationDestinationAddressPolicy(t *testing.T) {
blocked := []string{
"0.0.0.0", "10.0.0.1", "100.100.100.200", "127.0.0.1",
+75 -18
View File
@@ -46,10 +46,9 @@ func (s *Server) handleSMSContacts(w http.ResponseWriter, r *http.Request) {
}
deviceID := normalizeSMSDeviceFilter(r.URL.Query().Get("device_id"))
s.syncModemSMS(r.Context(), deviceID)
contacts, err := s.store.ListSMSContacts(r.Context(), store.SMSFilter{
DeviceID: deviceID,
Limit: queryLimit(r, 100),
})
filter := s.smsStoreFilter(r.Context(), deviceID, "")
filter.Limit = queryLimit(r, 100)
contacts, err := s.store.ListSMSContacts(r.Context(), filter)
if err != nil {
s.writeStoreError(w, err)
return
@@ -59,6 +58,7 @@ func (s *Server) handleSMSContacts(w http.ResponseWriter, r *http.Request) {
result = append(result, map[string]any{
"device_id": contact.DeviceID,
"device_name": contact.DeviceName,
"modem_imei": contact.ModemIMEI,
"imsi": contact.IMSI,
"local_phone": contact.LocalPhone,
"peer": contact.Peer,
@@ -78,6 +78,7 @@ func (s *Server) handleSMSContacts(w http.ResponseWriter, r *http.Request) {
func (s *Server) handleSMSThread(w http.ResponseWriter, r *http.Request) {
deviceID := normalizeSMSDeviceFilter(r.URL.Query().Get("device_id"))
modemIMEI := strings.TrimSpace(r.URL.Query().Get("modem_imei"))
imsi := strings.TrimSpace(r.URL.Query().Get("imsi"))
peer := strings.TrimSpace(r.URL.Query().Get("peer"))
if peer == "" {
@@ -87,12 +88,14 @@ func (s *Server) handleSMSThread(w http.ResponseWriter, r *http.Request) {
switch r.Method {
case http.MethodGet:
s.syncModemSMS(r.Context(), deviceID)
messages, err := s.store.ListSMSMessages(r.Context(), store.SMSFilter{
DeviceID: deviceID,
IMSI: imsi,
Peer: peer,
Limit: queryLimit(r, 100),
})
filter := s.smsStoreFilter(r.Context(), deviceID, modemIMEI)
filter.IMSI = imsi
filter.Peer = peer
filter.Limit = queryLimit(r, 100)
if beforeID, parseErr := strconv.ParseInt(strings.TrimSpace(r.URL.Query().Get("before_id")), 10, 64); parseErr == nil && beforeID > 0 {
filter.BeforeID = beforeID
}
messages, err := s.store.ListSMSMessages(r.Context(), filter)
if err != nil {
s.writeStoreError(w, err)
return
@@ -110,12 +113,11 @@ func (s *Server) handleSMSThread(w http.ResponseWriter, r *http.Request) {
}
writeJSON(w, http.StatusOK, map[string]any{"data": result})
case http.MethodDelete:
messages, err := s.store.ListSMSMessages(r.Context(), store.SMSFilter{
DeviceID: deviceID,
IMSI: imsi,
Peer: peer,
Limit: 1000,
})
filter := s.smsStoreFilter(r.Context(), deviceID, modemIMEI)
filter.IMSI = imsi
filter.Peer = peer
filter.Limit = 1000
messages, err := s.store.ListSMSMessages(r.Context(), filter)
if err != nil {
s.writeStoreError(w, err)
return
@@ -147,6 +149,35 @@ func normalizeSMSDeviceFilter(value string) string {
return value
}
// smsStoreFilter resolves a mutable configured device ID to the modem's stable
// IMEI. The ID is still used to address the live modem, but persisted history
// remains attached to the same hardware after the user renames that ID.
func (s *Server) smsStoreFilter(ctx context.Context, deviceID, requestedIMEI string) store.SMSFilter {
filter := store.SMSFilter{ModemIMEI: strings.TrimSpace(requestedIMEI)}
deviceID = strings.TrimSpace(deviceID)
if deviceID == "" {
return filter
}
filter.ModemIMEI = ""
filter.DeviceID = deviceID
config, err := s.store.Device(ctx, deviceID)
if err != nil {
return filter
}
imei := strings.TrimSpace(config.ModemIMEI)
if entry, _, present := s.physicalForConfig(config); present {
imei = firstNonEmpty(
snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMEI }),
imei,
)
}
if imei != "" {
filter.DeviceID = ""
filter.ModemIMEI = imei
}
return filter
}
// blockedSMSDestination reports whether the recipient is in a barred country.
// Normalization mirrors the PDU/IMS paths so the block cannot be sidestepped by
// dropping the leading "+" or using a 00 international prefix.
@@ -227,6 +258,10 @@ func (s *Server) handleSMSSend(w http.ResponseWriter, r *http.Request) {
return
}
imsi := snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMSI })
modemIMEI := firstNonEmpty(
snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMEI }),
config.ModemIMEI,
)
extra, _ := json.Marshal(map[string]any{
"encoding": result.Encoding,
"message_reference": result.MessageReference,
@@ -245,13 +280,14 @@ func (s *Server) handleSMSSend(w http.ResponseWriter, r *http.Request) {
})
messageID := fmt.Sprintf(
"at-submit:%s:%d:%d",
request.DeviceID,
firstNonEmpty(modemIMEI, request.DeviceID),
result.MessageReference,
result.SubmittedAt.UnixNano(),
)
saved, err := s.store.SaveSMSMessage(r.Context(), store.SMSMessage{
MessageID: messageID,
DeviceID: request.DeviceID,
ModemIMEI: modemIMEI,
IMSI: imsi,
Peer: result.To,
Direction: "outbound",
@@ -354,9 +390,14 @@ func (s *Server) writeIMSSMSSendResult(
"submission_status": result.SubmissionStatus,
})
imsi := snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMSI })
modemIMEI := snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMEI })
if config, configErr := s.store.Device(r.Context(), deviceID); configErr == nil {
modemIMEI = firstNonEmpty(modemIMEI, config.ModemIMEI)
}
saved, err := s.store.SaveSMSMessage(r.Context(), store.SMSMessage{
MessageID: fmt.Sprintf("ims-submit:%s:%d", deviceID, result.SubmittedAt.UnixNano()),
MessageID: fmt.Sprintf("ims-submit:%s:%d", firstNonEmpty(modemIMEI, deviceID), result.SubmittedAt.UnixNano()),
DeviceID: deviceID,
ModemIMEI: modemIMEI,
IMSI: imsi,
Peer: result.To,
Direction: "outbound",
@@ -494,11 +535,16 @@ func (s *Server) syncModemSMS(ctx context.Context, onlyDevice string) {
continue
}
imsi := snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMSI })
modemIMEI := firstNonEmpty(
snapshotString(entry.Snapshot, func(snapshot *device.Snapshot) string { return snapshot.IMEI }),
config.ModemIMEI,
)
for _, message := range messages {
if message.Direction == device.SMSDirectionStatusReport &&
message.MessageReference != nil && message.StatusCode != nil {
_, applyErr := s.store.ApplySMSDeliveryReport(ctx, store.SMSDeliveryReport{
DeviceID: config.ID,
ModemIMEI: modemIMEI,
IMSI: imsi,
Peer: message.To,
Source: "cellular_at",
@@ -535,6 +581,15 @@ func (s *Server) syncModemSMS(ctx context.Context, onlyDevice string) {
message.Index,
hex.EncodeToString(digest[:8]),
)
if message.Concat != nil && message.Concat.Total > 1 {
// A segment of a carrier-split long SMS. Address the whole message
// with a stable id so SaveSMSMessage folds every segment into one
// progressively merged row instead of one row per segment.
messageID = store.StableConcatMessageID(
"cellular_at", modemIMEI, config.ID, peer,
message.Concat.Reference, message.Concat.Total,
)
}
extra, _ := json.Marshal(map[string]any{
"modem_index": message.Index,
"storage": message.Storage,
@@ -550,6 +605,7 @@ func (s *Server) syncModemSMS(ctx context.Context, onlyDevice string) {
_, saveErr := s.store.SaveSMSMessage(ctx, store.SMSMessage{
MessageID: messageID,
DeviceID: config.ID,
ModemIMEI: modemIMEI,
IMSI: imsi,
Peer: peer,
Direction: direction,
@@ -605,6 +661,7 @@ func storedSMSResponse(message store.SMSMessage) map[string]any {
"id": message.ID,
"message_id": message.MessageID,
"device_id": message.DeviceID,
"modem_imei": message.ModemIMEI,
"imsi": message.IMSI,
"peer": message.Peer,
"direction": message.Direction,
+45 -3
View File
@@ -54,6 +54,48 @@ func TestSMSThreadAllDevicesUsesIMSIFilter(t *testing.T) {
}
}
func TestSMSThreadConfiguredDeviceUsesStableIMEI(t *testing.T) {
ctx := context.Background()
database, err := store.Open(ctx, ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
const imei = "867394042309830"
if err := database.UpsertDevice(ctx, store.Device{
ID: "ec20_2", Name: "EC20 renamed", ModemIMEI: imei,
}); err != nil {
t.Fatal(err)
}
if _, err := database.SaveSMSMessage(ctx, store.SMSMessage{
MessageID: "before-rename", DeviceID: "ec20_1", ModemIMEI: imei,
IMSI: "imsi-a", Peer: "VOXI", Direction: "inbound", Body: "history",
}); err != nil {
t.Fatal(err)
}
server := &Server{store: database}
request := httptest.NewRequest(
http.MethodGet,
"/api/sms/thread?device_id=ec20_2&imsi=imsi-a&peer=VOXI",
nil,
)
response := httptest.NewRecorder()
server.handleSMSThread(response, request)
if response.Code != http.StatusOK {
t.Fatalf("status = %d, body = %s", response.Code, response.Body.String())
}
var envelope struct {
Data []map[string]any `json:"data"`
}
if err := json.Unmarshal(response.Body.Bytes(), &envelope); err != nil {
t.Fatal(err)
}
if len(envelope.Data) != 1 || envelope.Data[0]["modem_imei"] != imei {
t.Fatalf("thread data = %#v", envelope.Data)
}
}
func TestNormalizeSMSDeviceFilter(t *testing.T) {
if got := normalizeSMSDeviceFilter(" ALL "); got != "" {
t.Fatalf("all filter = %q", got)
@@ -88,9 +130,9 @@ func TestSMSSendOutcome(t *testing.T) {
func TestBlockedSMSDestination(t *testing.T) {
tests := []struct {
name string
phone string
block bool
name string
phone string
block bool
}{
{"e164 china", "+8613800138000", true},
{"no plus china", "8613800138000", true},
+17 -3
View File
@@ -160,13 +160,27 @@ func validateSMSNotificationConfig(channel string, config map[string]any) error
func (s *Server) newSMSNotification(ctx context.Context, message store.SMSMessage) smsNotification {
name := ""
if device, err := s.store.Device(ctx, message.DeviceID); err == nil {
deviceID := message.DeviceID
if message.ModemIMEI != "" {
if devices, err := s.store.ListDevices(ctx); err == nil {
var newest time.Time
for _, candidate := range devices {
if candidate.ModemIMEI == message.ModemIMEI &&
(newest.IsZero() || candidate.UpdatedAt.After(newest)) {
deviceID = candidate.ID
name = strings.TrimSpace(candidate.Name)
newest = candidate.UpdatedAt
}
}
}
}
if device, err := s.store.Device(ctx, deviceID); err == nil {
name = strings.TrimSpace(device.Name)
}
return smsNotification{
DeviceID: message.DeviceID,
DeviceID: deviceID,
DeviceName: name,
DeviceLabel: firstNonEmpty(name, message.DeviceID, "--"),
DeviceLabel: firstNonEmpty(name, deviceID, "--"),
Number: firstNonEmpty(message.Peer, "--"),
Time: message.Timestamp,
Content: message.Body,
+178 -9
View File
@@ -11,6 +11,7 @@ import (
"io"
"net/http"
"net/http/httptest"
"regexp"
"strconv"
"strings"
"sync"
@@ -30,6 +31,8 @@ const (
telegramMaxDialDuration = 10 * time.Minute
)
var telegramTokenInURLPattern = regexp.MustCompile(`bot[0-9]{5,20}:[A-Za-z0-9_-]{20,128}`)
type telegramRuntimeConfig struct {
Token string
ChatID string
@@ -146,6 +149,9 @@ func (bot *telegramBot) poll(ctx context.Context) {
updates, pollErr := bot.getUpdates(pollContext, config, offset, 5)
cancel()
if pollErr != nil {
if ctx.Err() != nil {
return
}
bot.warn("poll Telegram updates", pollErr)
if !waitTelegram(ctx, telegramPollInterval) {
return
@@ -186,6 +192,10 @@ func (bot *telegramBot) bootstrap(ctx context.Context, config telegramRuntimeCon
{"command": "call", "description": "限时拨号并自动挂断(需要确认)"},
{"command": "calls", "description": "查看当前通话"},
{"command": "hangup", "description": "挂断通话"},
{"command": "at", "description": "向指定设备发送安全 AT 指令"},
{"command": "ussd", "description": "向指定设备发送 USSD 指令"},
{"command": "ussd_reply", "description": "回复交互式 USSD 会话"},
{"command": "ussd_cancel", "description": "取消交互式 USSD 会话"},
{"command": "help", "description": "查看命令帮助"},
}
_ = bot.call(requestContext, config, "setMyCommands", map[string]any{"commands": commands}, nil)
@@ -274,6 +284,34 @@ func (bot *telegramBot) handleUpdate(ctx context.Context, config telegramRuntime
bot.executeSimpleCallAction(ctx, config, message.Chat.ID, message.From.ID, strings.TrimSpace(remainder), "hangup")
case "calls":
bot.executeSimpleCallAction(ctx, config, message.Chat.ID, message.From.ID, strings.TrimSpace(remainder), "status")
case "at":
parts := splitTelegramArguments(remainder, 2)
if len(parts) != 2 {
bot.sendText(ctx, config, message.Chat.ID, "用法:/at <设备ID> <AT指令>\n示例:/at EC20 AT+CSQ", nil)
return
}
bot.handleATCommand(ctx, config, message.Chat.ID, message.From.ID, parts[0], parts[1])
case "ussd":
parts := strings.Fields(remainder)
if len(parts) != 2 {
bot.sendText(ctx, config, message.Chat.ID, "用法:/ussd <设备ID> <USSD代码>\n示例:/ussd EC20 *100#", nil)
return
}
bot.handleUSSDCommand(ctx, config, message.Chat.ID, message.From.ID, parts[0], parts[1])
case "ussd_reply":
parts := splitTelegramArguments(remainder, 2)
if len(parts) != 2 {
bot.sendText(ctx, config, message.Chat.ID, "用法:/ussd_reply <会话ID> <回复内容>", nil)
return
}
bot.handleUSSDReply(ctx, config, message.Chat.ID, message.From.ID, parts[0], parts[1])
case "ussd_cancel":
sessionID := strings.TrimSpace(remainder)
if sessionID == "" || strings.ContainsAny(sessionID, " \t\r\n") {
bot.sendText(ctx, config, message.Chat.ID, "用法:/ussd_cancel <会话ID>", nil)
return
}
bot.handleUSSDCancel(ctx, config, message.Chat.ID, message.From.ID, sessionID)
default:
bot.sendText(ctx, config, message.Chat.ID, "未知命令。发送 /help 查看可用操作。", nil)
}
@@ -329,6 +367,10 @@ func (bot *telegramBot) sendHelp(ctx context.Context, config telegramRuntimeConf
"/calls <设备ID> — 查看模块当前通话",
"/answer <设备ID> — 接听蜂窝来电",
"/hangup <设备ID> — 立即挂断",
"/at <设备ID> <AT指令> — 执行经过安全校验的单行 AT 指令",
"/ussd <设备ID> <代码> — 发送 USSD 指令",
"/ussd_reply <会话ID> <内容> — 回复交互式 USSD 菜单",
"/ussd_cancel <会话ID> — 取消交互式 USSD 会话",
"",
"Bot 不提供 eSIM 下载、删除或改名,也不采集或转发通话音频。控制命令只接受设置中的 Admin ID。",
}, "\n")
@@ -653,6 +695,112 @@ func (bot *telegramBot) executeSimpleCallAction(ctx context.Context, config tele
bot.server.recordAudit(ctx, fmt.Sprintf("telegram:%d", adminID), "telegram.call."+action, "device", deviceID, outcome, "telegram")
}
func (bot *telegramBot) handleATCommand(ctx context.Context, config telegramRuntimeConfig, chatID, adminID int64, deviceID, command string) {
result, err := bot.executeATCommand(ctx, deviceID, command)
outcome := "success"
if err != nil {
outcome = "failure"
bot.sendText(ctx, config, chatID, "AT 指令执行失败:"+err.Error(), nil)
} else {
bot.sendText(ctx, config, chatID, result, nil)
}
bot.server.recordAudit(ctx, fmt.Sprintf("telegram:%d", adminID), "telegram.at.execute", "device", deviceID, outcome, "telegram")
}
func (bot *telegramBot) executeATCommand(ctx context.Context, deviceID, command string) (string, error) {
command = strings.TrimSpace(command)
if err := validateATCommand(command); err != nil {
return "", err
}
_, _, physicalID, err := bot.device(deviceID)
if err != nil {
return "", err
}
operationContext, cancel := context.WithTimeout(ctx, 60*time.Second)
defer cancel()
response, err := bot.server.devices.ExecuteAT(operationContext, physicalID, command)
if err != nil {
return "", err
}
return fmt.Sprintf("设备:%s\n> %s\n\n%s", deviceID, command, formatTelegramAT(response)), nil
}
func (bot *telegramBot) handleUSSDCommand(ctx context.Context, config telegramRuntimeConfig, chatID, adminID int64, deviceID, code string) {
result, err := bot.executeUSSDCommand(ctx, deviceID, code)
outcome := "success"
if err != nil {
outcome = "failure"
bot.sendText(ctx, config, chatID, "USSD 指令执行失败:"+err.Error(), nil)
} else {
bot.sendText(ctx, config, chatID, formatTelegramUSSD(deviceID, result), nil)
}
bot.server.recordAudit(ctx, fmt.Sprintf("telegram:%d", adminID), "telegram.ussd.start", "device", deviceID, outcome, "telegram")
}
func (bot *telegramBot) executeUSSDCommand(ctx context.Context, deviceID, code string) (device.USSDResult, error) {
_, _, physicalID, err := bot.device(deviceID)
if err != nil {
return device.USSDResult{}, err
}
operationContext, cancel := context.WithTimeout(ctx, 90*time.Second)
defer cancel()
return bot.server.devices.USSD(operationContext, physicalID, strings.TrimSpace(code))
}
func (bot *telegramBot) handleUSSDReply(ctx context.Context, config telegramRuntimeConfig, chatID, adminID int64, sessionID, input string) {
operationContext, cancel := context.WithTimeout(ctx, 90*time.Second)
result, err := bot.server.devices.ContinueUSSD(operationContext, strings.TrimSpace(sessionID), strings.TrimSpace(input))
cancel()
outcome := "success"
if err != nil {
outcome = "failure"
bot.sendText(ctx, config, chatID, "USSD 回复失败:"+err.Error(), nil)
} else {
bot.sendText(ctx, config, chatID, formatTelegramUSSD("", result), nil)
}
bot.server.recordAudit(ctx, fmt.Sprintf("telegram:%d", adminID), "telegram.ussd.reply", "ussd_session", "interactive", outcome, "telegram")
}
func (bot *telegramBot) handleUSSDCancel(ctx context.Context, config telegramRuntimeConfig, chatID, adminID int64, sessionID string) {
operationContext, cancel := context.WithTimeout(ctx, 30*time.Second)
err := bot.server.devices.CancelUSSD(operationContext, strings.TrimSpace(sessionID))
cancel()
outcome := "success"
if err != nil {
outcome = "failure"
bot.sendText(ctx, config, chatID, "取消 USSD 会话失败:"+err.Error(), nil)
} else {
bot.sendText(ctx, config, chatID, "USSD 会话已取消。", nil)
}
bot.server.recordAudit(ctx, fmt.Sprintf("telegram:%d", adminID), "telegram.ussd.cancel", "ussd_session", "interactive", outcome, "telegram")
}
func formatTelegramUSSD(deviceID string, result device.USSDResult) string {
lines := make([]string, 0, 7)
if strings.TrimSpace(deviceID) != "" {
lines = append(lines, "设备:"+strings.TrimSpace(deviceID))
}
if strings.TrimSpace(result.Code) != "" {
lines = append(lines, "USSD"+strings.TrimSpace(result.Code))
}
lines = append(lines, "状态:"+firstNonEmpty(strings.TrimSpace(result.Status), "final"))
if strings.TrimSpace(result.Text) != "" {
lines = append(lines, "\n"+strings.TrimSpace(result.Text))
} else if strings.TrimSpace(result.Raw) != "" {
lines = append(lines, "\n"+strings.TrimSpace(result.Raw))
} else {
lines = append(lines, "\n网络未返回文本内容。")
}
if result.Continueable && strings.TrimSpace(result.SessionID) != "" {
lines = append(lines,
"\n网络正在等待输入。",
"回复:/ussd_reply "+result.SessionID+" <内容>",
"取消:/ussd_cancel "+result.SessionID,
)
}
return strings.Join(lines, "\n")
}
func (bot *telegramBot) handleVoWiFi(ctx context.Context, config telegramRuntimeConfig, chatID, adminID int64, deviceID, operation string) {
stored, entry, _, err := bot.device(deviceID)
if err != nil {
@@ -754,6 +902,15 @@ func (bot *telegramBot) notifyInboundSMS(ctx context.Context) {
bot.warn("list Telegram SMS notifications", listErr)
} else {
for _, message := range messages {
if !store.ConcatSMSReadyToNotify(message.MessageID, message.Extra) {
// A carrier-split long SMS still waiting for segments. Hold
// the notification but advance the cursor so the partial row
// is not reconsidered every poll; when the final segment
// merges, the row re-enters with a fresh id and is pushed
// here as one complete message.
cursor = message.ID
continue
}
text := fmt.Sprintf("📩 新短信\n设备:%s\n来自:%s\n时间:%s\n\n%s", message.DeviceID, message.Peer, message.Timestamp.Local().Format("2006-01-02 15:04:05"), message.Body)
if sendErr := bot.sendText(ctx, config, 0, text, nil); sendErr != nil {
bot.warn("send Telegram SMS notification", sendErr)
@@ -811,7 +968,7 @@ func (bot *telegramBot) loadConfig(ctx context.Context) (telegramRuntimeConfig,
Proxy: configString(raw, "proxy"),
}
if config.BaseURL == "" {
config.BaseURL = "https://api.telegram.org"
config.BaseURL = defaultTelegramBaseURL
}
if admin := configString(raw, "admin_id"); admin != "" {
config.AdminID, err = strconv.ParseInt(admin, 10, 64)
@@ -826,12 +983,10 @@ func (bot *telegramBot) loadConfig(ctx context.Context) (telegramRuntimeConfig,
}
func (bot *telegramBot) call(ctx context.Context, config telegramRuntimeConfig, method string, payload any, result any) error {
base, err := validateOutboundURL(ctx, config.BaseURL, true)
base, err := validateTelegramAPIURL(ctx, config.BaseURL, config.Token, method)
if err != nil {
return err
return redactTelegramError(err, config.Token)
}
base.Path = strings.TrimRight(base.Path, "/") + "/bot" + config.Token + "/" + method
base.RawPath, base.RawQuery, base.Fragment = "", "", ""
body, err := json.Marshal(payload)
if err != nil {
return err
@@ -842,13 +997,13 @@ func (bot *telegramBot) call(ctx context.Context, config telegramRuntimeConfig,
}
request, err := http.NewRequestWithContext(ctx, http.MethodPost, base.String(), bytes.NewReader(body))
if err != nil {
return err
return redactTelegramError(err, config.Token)
}
request.Header.Set("Content-Type", "application/json")
request.Header.Set("User-Agent", "vocat-telegram-bot/1")
response, err := client.Do(request)
if err != nil {
return err
return redactTelegramError(err, config.Token)
}
defer response.Body.Close()
responseBody, err := io.ReadAll(io.LimitReader(response.Body, 2<<20))
@@ -902,7 +1057,7 @@ func (bot *telegramBot) warn(message string, err error) {
return
}
now := time.Now()
text := err.Error()
text := redactTelegramText(err.Error(), "")
bot.logMu.Lock()
if text == bot.lastLogText && now.Sub(bot.lastLogTime) < time.Minute {
bot.logMu.Unlock()
@@ -910,7 +1065,21 @@ func (bot *telegramBot) warn(message string, err error) {
}
bot.lastLogText, bot.lastLogTime = text, now
bot.logMu.Unlock()
bot.server.logger.Warn(message, "error", err)
bot.server.logger.Warn(message, "error", text)
}
func redactTelegramError(err error, token string) error {
if err == nil {
return nil
}
return errors.New(redactTelegramText(err.Error(), token))
}
func redactTelegramText(value, token string) string {
if strings.TrimSpace(token) != "" {
value = strings.ReplaceAll(value, token, "[REDACTED]")
}
return telegramTokenInURLPattern.ReplaceAllString(value, "bot[REDACTED]")
}
func parseTelegramCommand(text string) (string, string) {
+121
View File
@@ -1,12 +1,61 @@
package server
import (
"context"
"errors"
"strings"
"testing"
"time"
"vocat/internal/device"
"vocat/internal/modem"
"vocat/internal/store"
)
func TestTelegramAPIURLSupportsBaseAndTemplate(t *testing.T) {
tests := []struct {
name string
baseURL string
want string
}{
{
name: "base URL",
baseURL: "https://api.telegram.org",
want: "https://api.telegram.org/bot123456:test-token/sendMessage",
},
{
name: "reverse proxy template",
baseURL: "https://telegram.example.com/bot%s/%s",
want: "https://telegram.example.com/bot123456:test-token/sendMessage",
},
}
for _, item := range tests {
t.Run(item.name, func(t *testing.T) {
got, err := telegramAPIURL(item.baseURL, "123456:test-token", "sendMessage")
if err != nil {
t.Fatal(err)
}
if got.String() != item.want {
t.Fatalf("telegramAPIURL() = %q, want %q", got, item.want)
}
})
}
}
func TestTelegramAPIURLRejectsMalformedTemplates(t *testing.T) {
for _, value := range []string{
"https://telegram.example.com/bot%s/sendMessage",
"https://%s.example.com/bot/token/%s",
"http://telegram.example.com/bot%s/%s",
} {
if _, err := telegramAPIURL(value, "123456:test-token", "sendMessage"); err == nil {
t.Errorf("telegramAPIURL(%q) unexpectedly succeeded", value)
} else if strings.TrimSpace(err.Error()) == "" {
t.Errorf("telegramAPIURL(%q) returned an empty error", value)
}
}
}
func TestParseTelegramCommand(t *testing.T) {
command, remainder := parseTelegramCommand(" /sms@vocat_bot EC20 +447700900123 hello world ")
if command != "sms" || remainder != "EC20 +447700900123 hello world" {
@@ -60,3 +109,75 @@ func TestFormatTelegramATIncludesFinalResult(t *testing.T) {
t.Fatalf("formatTelegramAT(lines) = %q", got)
}
}
func TestTelegramExecutesGuardedATForConfiguredDevice(t *testing.T) {
database, err := store.Open(context.Background(), ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
if err := database.UpsertDevice(context.Background(), store.Device{ID: "EC20", Name: "EC20"}); err != nil {
t.Fatal(err)
}
bot := &telegramBot{server: &Server{
store: database,
devices: fakeDeviceController{
entry: device.Device{ID: "EC20", Discovered: true},
atResponse: modem.Response{Lines: []string{"+CSQ: 18,99"}, Final: "OK"},
},
}}
result, err := bot.executeATCommand(context.Background(), "EC20", "AT+CSQ")
if err != nil {
t.Fatal(err)
}
for _, expected := range []string{"设备:EC20", "> AT+CSQ", "+CSQ: 18,99", "OK"} {
if !strings.Contains(result, expected) {
t.Fatalf("AT result %q does not contain %q", result, expected)
}
}
if _, err := bot.executeATCommand(context.Background(), "EC20", "AT+CFUN=0"); err == nil {
t.Fatal("guarded AT command unexpectedly succeeded")
}
}
func TestTelegramExecutesInteractiveUSSDForConfiguredDevice(t *testing.T) {
database, err := store.Open(context.Background(), ":memory:")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = database.Close() })
if err := database.UpsertDevice(context.Background(), store.Device{ID: "EC20", Name: "EC20"}); err != nil {
t.Fatal(err)
}
bot := &telegramBot{server: &Server{
store: database,
devices: fakeDeviceController{
entry: device.Device{ID: "EC20", Discovered: true},
ussdResult: device.USSDResult{
Code: "*100#", Text: "1. Balance\n2. Bundles", Status: "awaiting_input",
SessionID: "0123456789abcdef", Continueable: true,
},
},
}}
result, err := bot.executeUSSDCommand(context.Background(), "EC20", "*100#")
if err != nil {
t.Fatal(err)
}
formatted := formatTelegramUSSD("EC20", result)
for _, expected := range []string{
"设备:EC20", "状态:awaiting_input", "1. Balance", "/ussd_reply 0123456789abcdef", "/ussd_cancel 0123456789abcdef",
} {
if !strings.Contains(formatted, expected) {
t.Fatalf("USSD result %q does not contain %q", formatted, expected)
}
}
}
func TestTelegramErrorsRedactBotTokens(t *testing.T) {
token := "1234567890:abcdefghijklmnopqrstuvwxyzABCDE"
err := errors.New(`Post "https://api.telegram.org/bot` + token + `/getUpdates": context canceled`)
redacted := redactTelegramError(err, token)
if strings.Contains(redacted.Error(), token) || !strings.Contains(redacted.Error(), "bot[REDACTED]") {
t.Fatalf("redacted error = %q", redacted)
}
}
+65
View File
@@ -0,0 +1,65 @@
package server
import (
"context"
"errors"
"net/url"
"strings"
)
const defaultTelegramBaseURL = "https://api.telegram.org"
// telegramAPIURL accepts either a Telegram API base URL or a printf-style
// endpoint template whose two %s placeholders are the bot token and method.
func telegramAPIURL(baseURL, token, method string) (*url.URL, error) {
raw := strings.TrimSpace(baseURL)
if raw == "" {
raw = defaultTelegramBaseURL
}
placeholderCount := strings.Count(raw, "%s")
if placeholderCount != 0 && placeholderCount != 2 {
return nil, errors.New("Telegram API URL must contain either no %s placeholders or exactly two")
}
if placeholderCount == 2 {
if telegramPlaceholderInAuthority(raw) {
return nil, errors.New("Telegram API URL placeholders are not allowed in the host")
}
endpoint := strings.Replace(raw, "%s", token, 1)
endpoint = strings.Replace(endpoint, "%s", method, 1)
return parseOutboundURL(endpoint, true)
}
parsed, err := parseOutboundURL(raw, true)
if err != nil {
return nil, err
}
parsed.Path = strings.TrimRight(parsed.Path, "/") + "/bot" + token + "/" + method
parsed.RawPath = ""
parsed.RawQuery = ""
parsed.Fragment = ""
return parsed, nil
}
func validateTelegramAPIURL(ctx context.Context, baseURL, token, method string) (*url.URL, error) {
parsed, err := telegramAPIURL(baseURL, token, method)
if err != nil {
return nil, err
}
if _, err := resolvePublicAddresses(ctx, parsed.Hostname()); err != nil {
return nil, err
}
return parsed, nil
}
func telegramPlaceholderInAuthority(raw string) bool {
schemeEnd := strings.Index(raw, "://")
if schemeEnd < 0 {
return false
}
authority := raw[schemeEnd+3:]
if end := strings.IndexAny(authority, "/?#"); end >= 0 {
authority = authority[:end]
}
return strings.Contains(authority, "%s")
}
+75 -7
View File
@@ -13,6 +13,27 @@ type contextExecer interface {
ExecContext(context.Context, string, ...any) (sql.Result, error)
}
const (
DeviceTypeWiFi410 = "wifi_410"
DeviceTypeDJI4G = "dji_4g"
DeviceTypePCIeEC20EC25 = "pcie_ec20_ec25"
)
// NormalizeDeviceType returns a stable persisted device type identifier.
// Empty values use the legacy EC20/EC25 type for backwards compatibility.
func NormalizeDeviceType(value string) string {
switch strings.ToLower(strings.TrimSpace(value)) {
case DeviceTypeWiFi410:
return DeviceTypeWiFi410
case DeviceTypeDJI4G:
return DeviceTypeDJI4G
case "", DeviceTypePCIeEC20EC25:
return DeviceTypePCIeEC20EC25
default:
return ""
}
}
func (s *Store) UpsertDevice(ctx context.Context, value Device) error {
return upsertDevice(ctx, s.db, value)
}
@@ -73,6 +94,10 @@ func upsertDevice(ctx context.Context, executor contextExecer, value Device) err
if value.Name == "" {
return errors.New("device name is required")
}
value.DeviceType = NormalizeDeviceType(value.DeviceType)
if value.DeviceType == "" {
return errors.New("unsupported device type")
}
if value.ProxyPort < 0 || value.ProxyPort > 65535 {
return errors.New("device proxy port must be between 0 and 65535")
}
@@ -133,15 +158,16 @@ func upsertDevice(ctx context.Context, executor contextExecer, value Device) err
_, err = executor.ExecContext(ctx, `
INSERT INTO devices (
id, name, interface, control_device, at_port, usb_path,
id, name, device_type, interface, control_device, at_port, usb_path,
audio_device, modem_imei, apn, proxy_port, baud_rate,
data_bits, stop_bits, parity, device_backend, esim_transport,
qmi_use_proxy, qmi_proxy_path, qmi_proxy_executable,
network_enabled, sms_enabled, vowifi_enabled, extra_json,
created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(id) DO UPDATE SET
name = excluded.name,
device_type = excluded.device_type,
interface = excluded.interface,
control_device = excluded.control_device,
at_port = excluded.at_port,
@@ -165,7 +191,7 @@ func upsertDevice(ctx context.Context, executor contextExecer, value Device) err
extra_json = excluded.extra_json,
updated_at = excluded.updated_at
`,
value.ID, value.Name, value.Interface, value.ControlDevice, value.ATPort,
value.ID, value.Name, value.DeviceType, value.Interface, value.ControlDevice, value.ATPort,
value.USBPath, value.AudioDevice, value.ModemIMEI, value.APN,
value.ProxyPort, value.BaudRate, value.DataBits, value.StopBits,
value.Parity, value.DeviceBackend, value.ESIMTransport,
@@ -206,15 +232,56 @@ func (s *Store) ListDevices(ctx context.Context) ([]Device, error) {
}
func (s *Store) DeleteDevice(ctx context.Context, id string) error {
result, err := s.db.ExecContext(ctx, `DELETE FROM devices WHERE id = ?`, id)
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return fmt.Errorf("begin delete device %q: %w", id, err)
}
defer tx.Rollback()
var modemIMEI string
if err := tx.QueryRowContext(ctx, `SELECT modem_imei FROM devices WHERE id = ?`, id).Scan(&modemIMEI); err != nil {
if errors.Is(err, sql.ErrNoRows) {
return ErrNotFound
}
return fmt.Errorf("read device %q before deletion: %w", id, err)
}
// SMS history must outlive a mutable configured device ID. Anchor any
// legacy ID-owned rows to the physical modem before the device row and its
// runtime records are removed.
if strings.TrimSpace(modemIMEI) != "" {
if _, err := tx.ExecContext(ctx, `
DELETE FROM sms_messages AS legacy
WHERE legacy.device_id = ? AND legacy.modem_imei = '' AND legacy.message_id <> ''
AND EXISTS (
SELECT 1 FROM sms_messages current
WHERE current.modem_imei = ?
AND current.message_id = legacy.message_id
)
`, id, strings.TrimSpace(modemIMEI)); err != nil {
return fmt.Errorf("deduplicate SMS history for device %q: %w", id, err)
}
if _, err := tx.ExecContext(ctx, `
UPDATE sms_messages
SET modem_imei = ?, updated_at = ?
WHERE device_id = ? AND modem_imei = ''
`, strings.TrimSpace(modemIMEI), time.Now().UTC().Unix(), id); err != nil {
return fmt.Errorf("anchor SMS history for device %q: %w", id, err)
}
}
result, err := tx.ExecContext(ctx, `DELETE FROM devices WHERE id = ?`, id)
if err != nil {
return fmt.Errorf("delete device %q: %w", id, err)
}
return requireAffected(result)
if err := requireAffected(result); err != nil {
return err
}
if err := tx.Commit(); err != nil {
return fmt.Errorf("commit delete device %q: %w", id, err)
}
return nil
}
const deviceSelect = `
SELECT id, name, interface, control_device, at_port, usb_path,
SELECT id, name, device_type, interface, control_device, at_port, usb_path,
audio_device, modem_imei, apn, proxy_port, baud_rate, data_bits,
stop_bits, parity, device_backend, esim_transport, qmi_use_proxy,
qmi_proxy_path, qmi_proxy_executable, network_enabled, sms_enabled,
@@ -227,7 +294,7 @@ func scanDevice(row rowScanner) (Device, error) {
var extra string
var createdAt, updatedAt int64
err := row.Scan(
&value.ID, &value.Name, &value.Interface, &value.ControlDevice,
&value.ID, &value.Name, &value.DeviceType, &value.Interface, &value.ControlDevice,
&value.ATPort, &value.USBPath, &value.AudioDevice, &value.ModemIMEI,
&value.APN, &value.ProxyPort, &value.BaudRate, &value.DataBits,
&value.StopBits, &value.Parity, &value.DeviceBackend,
@@ -245,6 +312,7 @@ func scanDevice(row rowScanner) (Device, error) {
value.NetworkEnabled = networkEnabled != 0
value.SMSEnabled = smsEnabled != 0
value.VoWiFiEnabled = vowifiEnabled != 0
value.DeviceType = NormalizeDeviceType(value.DeviceType)
value.Extra = []byte(extra)
value.CreatedAt = time.Unix(createdAt, 0).UTC()
value.UpdatedAt = time.Unix(updatedAt, 0).UTC()
+134 -1
View File
@@ -70,6 +70,76 @@ func TestMigrationFromAuthenticationSchema(t *testing.T) {
}
}
func TestMigration7BackfillsSMSModemIMEI(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "sms-imei.db")
raw, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
for version := 1; version <= 6; version++ {
for _, statement := range migrationStatements(version) {
if _, err := raw.ExecContext(ctx, statement); err != nil {
t.Fatalf("create v%d schema: %v", version, err)
}
}
}
if _, err := raw.ExecContext(ctx, `
INSERT INTO devices (id, name, modem_imei, created_at, updated_at)
VALUES ('ec20_1', 'EC20', '867394042309830', 100, 100);
INSERT INTO sms_messages (
message_id, device_id, peer, direction, message_time, created_at, updated_at
) VALUES ('legacy-message', 'ec20_1', 'VOXI', 'inbound', 100, 100, 100);
PRAGMA user_version = 6;
`); err != nil {
t.Fatal(err)
}
if err := raw.Close(); err != nil {
t.Fatal(err)
}
database := openTestStore(t, path)
messages, err := database.ListSMSMessages(ctx, SMSFilter{ModemIMEI: "867394042309830"})
if err != nil || len(messages) != 1 || messages[0].DeviceID != "ec20_1" {
t.Fatalf("migrated SMS = %#v, %v", messages, err)
}
}
func TestMigration8DefaultsExistingDevicesToPCIeType(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "device-type.db")
raw, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
for version := 1; version <= 7; version++ {
for _, statement := range migrationStatements(version) {
if _, err := raw.ExecContext(ctx, statement); err != nil {
t.Fatalf("create v%d schema: %v", version, err)
}
}
}
if _, err := raw.ExecContext(ctx, `
INSERT INTO devices (id, name, created_at, updated_at)
VALUES ('legacy', 'Legacy modem', 100, 100);
PRAGMA user_version = 7;
`); err != nil {
t.Fatal(err)
}
if err := raw.Close(); err != nil {
t.Fatal(err)
}
database := openTestStore(t, path)
got, err := database.Device(ctx, "legacy")
if err != nil {
t.Fatal(err)
}
if got.DeviceType != DeviceTypePCIeEC20EC25 {
t.Fatalf("legacy device type = %q", got.DeviceType)
}
}
func TestMigration4PreservesIMSRedeliveryAndUsesReceiptTime(t *testing.T) {
ctx := context.Background()
path := filepath.Join(t.TempDir(), "ims-redelivery.db")
@@ -135,6 +205,7 @@ func TestDeviceStateRoundTripAndCascade(t *testing.T) {
device := Device{
ID: "ec20-1",
Name: "EC20 一号",
DeviceType: DeviceTypeDJI4G,
Interface: "wwan0",
ControlDevice: "/dev/cdc-wdm0",
ATPort: "/dev/ttyUSB2",
@@ -184,7 +255,7 @@ func TestDeviceStateRoundTripAndCascade(t *testing.T) {
if err != nil {
t.Fatal(err)
}
if gotDevice.BaudRate != 115200 || gotDevice.DataBits != 8 ||
if gotDevice.DeviceType != DeviceTypeDJI4G || gotDevice.BaudRate != 115200 || gotDevice.DataBits != 8 ||
gotDevice.StopBits != 1 || gotDevice.DeviceBackend != "at" {
t.Fatalf("device defaults not applied: %+v", gotDevice)
}
@@ -291,6 +362,68 @@ func TestSMSPersistenceAndDerivedThreads(t *testing.T) {
}
}
func TestSMSHistoryFollowsModemIMEIAfterDeviceIDRename(t *testing.T) {
ctx := context.Background()
database := openTestStore(t, ":memory:")
const imei = "867394042309830"
if err := database.UpsertDevice(ctx, Device{
ID: "ec20_1", Name: "EC20 old", ModemIMEI: imei,
}); err != nil {
t.Fatal(err)
}
if _, err := database.SaveSMSMessage(ctx, SMSMessage{
MessageID: "network-old", DeviceID: "ec20_1",
IMSI: "23415", Peer: "VOXI", Direction: "inbound", Body: "before rename",
Timestamp: time.Unix(1_700_000_000, 0).UTC(),
}); err != nil {
t.Fatal(err)
}
if err := database.DeleteDevice(ctx, "ec20_1"); err != nil {
t.Fatal(err)
}
if err := database.UpsertDevice(ctx, Device{
ID: "ec20_2", Name: "EC20 renamed", ModemIMEI: imei,
}); err != nil {
t.Fatal(err)
}
if _, err := database.SaveSMSMessage(ctx, SMSMessage{
MessageID: "network-new", DeviceID: "ec20_2", ModemIMEI: imei,
IMSI: "23415", Peer: "VOXI", Direction: "inbound", Body: "after rename",
Timestamp: time.Unix(1_700_000_060, 0).UTC(),
}); err != nil {
t.Fatal(err)
}
contacts, err := database.ListSMSContacts(ctx, SMSFilter{ModemIMEI: imei})
if err != nil {
t.Fatal(err)
}
if len(contacts) != 1 || contacts[0].DeviceID != "ec20_2" ||
contacts[0].ModemIMEI != imei || contacts[0].MessageCount != 2 {
t.Fatalf("renamed hardware contact = %#v", contacts)
}
messages, err := database.ListSMSMessages(ctx, SMSFilter{
ModemIMEI: imei, IMSI: "23415", Peer: "VOXI",
})
if err != nil || len(messages) != 2 {
t.Fatalf("renamed hardware messages = %#v, %v", messages, err)
}
// A retry that arrives after the rename updates the same hardware message,
// rather than duplicating it under the new configured ID.
if _, err := database.SaveSMSMessage(ctx, SMSMessage{
MessageID: "network-old", DeviceID: "ec20_2", ModemIMEI: imei,
IMSI: "23415", Peer: "VOXI", Direction: "inbound", Body: "retry",
Timestamp: time.Unix(1_700_000_000, 0).UTC(),
}); err != nil {
t.Fatal(err)
}
messages, err = database.ListSMSMessages(ctx, SMSFilter{ModemIMEI: imei})
if err != nil || len(messages) != 2 {
t.Fatalf("retry after rename messages = %#v, %v", messages, err)
}
}
func TestListInboundSMSAfterIDUsesDurableInsertionCursor(t *testing.T) {
ctx := context.Background()
database := openTestStore(t, ":memory:")
+30
View File
@@ -81,6 +81,36 @@ func migrationStatements(version int) []string {
`CREATE INDEX IF NOT EXISTS device_proxy_bindings_proxy_idx
ON device_proxy_bindings(upstream_proxy_id)`,
}
case 7:
return []string{
`ALTER TABLE sms_messages
ADD COLUMN modem_imei TEXT NOT NULL DEFAULT ''`,
`UPDATE sms_messages
SET modem_imei = COALESCE((
SELECT NULLIF(d.modem_imei, '')
FROM devices d
WHERE d.id = sms_messages.device_id
), '')
WHERE modem_imei = ''`,
`DELETE FROM sms_messages
WHERE modem_imei <> '' AND message_id <> ''
AND id NOT IN (
SELECT MIN(id)
FROM sms_messages
WHERE modem_imei <> '' AND message_id <> ''
GROUP BY modem_imei, message_id
)`,
`CREATE UNIQUE INDEX IF NOT EXISTS sms_messages_hardware_external_id_idx
ON sms_messages(modem_imei, message_id)
WHERE modem_imei <> '' AND message_id <> ''`,
`CREATE INDEX IF NOT EXISTS sms_messages_hardware_thread_idx
ON sms_messages(modem_imei, imsi, peer, message_time DESC, id DESC)`,
}
case 8:
return []string{
`ALTER TABLE devices
ADD COLUMN device_type TEXT NOT NULL DEFAULT 'pcie_ec20_ec25'`,
}
default:
return nil
}
+12 -7
View File
@@ -17,6 +17,7 @@ const SecretMask = "********"
type Device struct {
ID string
Name string
DeviceType string
Interface string
ControlDevice string
ATPort string
@@ -127,6 +128,7 @@ type SMSMessage struct {
ID int64
MessageID string
DeviceID string
ModemIMEI string
IMSI string
Peer string
Direction string
@@ -143,19 +145,21 @@ type SMSMessage struct {
}
type SMSFilter struct {
DeviceID string
IMSI string
Peer string
Since time.Time
Until time.Time
BeforeID int64
Limit int
DeviceID string
ModemIMEI string
IMSI string
Peer string
Since time.Time
Until time.Time
BeforeID int64
Limit int
}
// SMSDeliveryReport is network evidence for one submitted SMS part. The
// message reference is the TP-MR returned in SMS-STATUS-REPORT.
type SMSDeliveryReport struct {
DeviceID string
ModemIMEI string
IMSI string
Peer string
Source string
@@ -170,6 +174,7 @@ type SMSDeliveryReport struct {
type SMSContact struct {
DeviceID string
DeviceName string
ModemIMEI string
IMSI string
LocalPhone string
Peer string
+113 -25
View File
@@ -40,6 +40,7 @@ func saveSMSMessage(
value SMSMessage,
) (SMSMessage, error) {
value.DeviceID = strings.TrimSpace(value.DeviceID)
value.ModemIMEI = strings.TrimSpace(value.ModemIMEI)
value.Peer = strings.TrimSpace(value.Peer)
value.Direction = strings.ToLower(strings.TrimSpace(value.Direction))
if value.DeviceID == "" {
@@ -64,6 +65,58 @@ func saveSMSMessage(
return SMSMessage{}, fmt.Errorf("normalize SMS extra data: %w", err)
}
now := time.Now().UTC()
// Concatenated (long) SMS arrive as one segment per delivery. Ingest points
// address the whole message with a stable "concat:" message id and carry the
// segment text plus its UDH sequence in Extra. Fold each segment into a single
// stored row so history, the web thread, and Telegram show one progressive
// message that fills in as the remaining segments arrive.
if isConcatSMSMessageID(value.MessageID) {
hardwareKey := smsHardwareKey(value.ModemIMEI, value.DeviceID)
existing, existingErr := scanSMSMessage(executor.QueryRowContext(
ctx,
smsMessageSelect+` WHERE
COALESCE(NULLIF(modem_imei, ''), 'device:' || device_id) = ?
AND message_id = ?`,
hardwareKey,
value.MessageID,
))
if existingErr != nil && !errors.Is(existingErr, ErrNotFound) {
return SMSMessage{}, fmt.Errorf("read existing concatenated SMS: %w", existingErr)
}
var existingExtra json.RawMessage
if existingErr == nil {
existingExtra = existing.Extra
}
mergedBody, mergedExtra, changed, mergeErr := mergeConcatSegment(existingExtra, value.Body, extra)
if mergeErr != nil {
return SMSMessage{}, fmt.Errorf("merge concatenated SMS segment: %w", mergeErr)
}
if existingErr == nil && !changed {
// This segment is already folded into the stored row (a periodic modem
// rescan redelivers every segment). Leave the row untouched so the
// durable id stays put and Telegram does not re-notify.
return existing, nil
}
value.Body = mergedBody
extra = mergedExtra
if existingErr == nil {
// A new segment advanced the message. Replace the stale partial row so
// the merged row receives a fresh durable id; the Telegram id-cursor
// then surfaces the now-more-complete message exactly once. Carry
// forward identity and history fields.
if _, delErr := executor.ExecContext(ctx, `DELETE FROM sms_messages WHERE id = ?`, existing.ID); delErr != nil {
return SMSMessage{}, fmt.Errorf("replace concatenated SMS: %w", delErr)
}
value.ID = 0
value.CreatedAt = existing.CreatedAt
value.Read = value.Read || existing.Read
if !existing.Timestamp.IsZero() &&
(value.Timestamp.IsZero() || existing.Timestamp.Before(value.Timestamp)) {
value.Timestamp = existing.Timestamp
}
}
}
if value.Timestamp.IsZero() {
value.Timestamp = now
}
@@ -77,13 +130,13 @@ func saveSMSMessage(
if value.ID > 0 {
result, err := executor.ExecContext(ctx, `
UPDATE sms_messages SET
message_id = ?, device_id = ?, imsi = ?, peer = ?,
message_id = ?, device_id = ?, modem_imei = ?, imsi = ?, peer = ?,
direction = ?, body = ?, message_time = ?, status = ?,
source = ?, parts_total = ?, delivery_state = ?, is_read = ?,
extra_json = ?, updated_at = ?
WHERE id = ?
`,
value.MessageID, value.DeviceID, value.IMSI, value.Peer,
value.MessageID, value.DeviceID, value.ModemIMEI, value.IMSI, value.Peer,
value.Direction, value.Body, value.Timestamp.Unix(), value.Status,
value.Source, value.PartsTotal, value.DeliveryState,
boolInt(value.Read), string(extra), value.UpdatedAt.Unix(), value.ID,
@@ -99,11 +152,16 @@ func saveSMSMessage(
result, err := executor.ExecContext(ctx, `
INSERT INTO sms_messages (
message_id, device_id, imsi, peer, direction, body, message_time,
message_id, device_id, modem_imei, imsi, peer, direction, body, message_time,
status, source, parts_total, delivery_state, is_read, extra_json,
created_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(device_id, message_id) WHERE message_id <> '' DO UPDATE SET
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT DO UPDATE SET
device_id = excluded.device_id,
modem_imei = CASE
WHEN excluded.modem_imei <> '' THEN excluded.modem_imei
ELSE sms_messages.modem_imei
END,
imsi = excluded.imsi,
peer = excluded.peer,
direction = excluded.direction,
@@ -117,7 +175,7 @@ func saveSMSMessage(
extra_json = excluded.extra_json,
updated_at = excluded.updated_at
`,
value.MessageID, value.DeviceID, value.IMSI, value.Peer,
value.MessageID, value.DeviceID, value.ModemIMEI, value.IMSI, value.Peer,
value.Direction, value.Body, value.Timestamp.Unix(), value.Status,
value.Source, value.PartsTotal, value.DeliveryState,
boolInt(value.Read), string(extra), value.CreatedAt.Unix(),
@@ -127,10 +185,13 @@ func saveSMSMessage(
return SMSMessage{}, fmt.Errorf("save SMS: %w", err)
}
if value.MessageID != "" {
hardwareKey := smsHardwareKey(value.ModemIMEI, value.DeviceID)
return scanSMSMessage(executor.QueryRowContext(
ctx,
smsMessageSelect+` WHERE device_id = ? AND message_id = ?`,
value.DeviceID,
smsMessageSelect+` WHERE
COALESCE(NULLIF(modem_imei, ''), 'device:' || device_id) = ?
AND message_id = ?`,
hardwareKey,
value.MessageID,
))
}
@@ -207,7 +268,9 @@ func (s *Store) ListInboundSMSAfterID(ctx context.Context, afterID int64, limit
// outbound submission and advances its aggregate delivery state. Multipart
// messages become delivered only after every submitted part is reported.
func (s *Store) ApplySMSDeliveryReport(ctx context.Context, report SMSDeliveryReport) (SMSMessage, error) {
if report.DeviceID == "" || report.MessageReference < 0 || report.MessageReference > 255 {
report.DeviceID = strings.TrimSpace(report.DeviceID)
report.ModemIMEI = strings.TrimSpace(report.ModemIMEI)
if (report.DeviceID == "" && report.ModemIMEI == "") || report.MessageReference < 0 || report.MessageReference > 255 {
return SMSMessage{}, errors.New("invalid SMS delivery report identity")
}
if report.ReceivedAt.IsZero() {
@@ -219,7 +282,7 @@ func (s *Store) ApplySMSDeliveryReport(ctx context.Context, report SMSDeliveryRe
}
defer tx.Rollback()
query := smsMessageSelect + `
WHERE device_id = ?
WHERE ((? <> '' AND modem_imei = ?) OR (? = '' AND device_id = ?))
AND direction IN ('outbound', 'sent')
AND (? = '' OR imsi = ?)
AND (? = '' OR peer = ?)
@@ -229,7 +292,7 @@ func (s *Store) ApplySMSDeliveryReport(ctx context.Context, report SMSDeliveryRe
rows, err := tx.QueryContext(
ctx,
query,
report.DeviceID,
report.ModemIMEI, report.ModemIMEI, report.ModemIMEI, report.DeviceID,
report.IMSI, report.IMSI,
report.Peer, report.Peer,
report.Source, report.Source,
@@ -449,28 +512,42 @@ func (s *Store) MarkSMSThreadRead(
func (s *Store) ListSMSContacts(ctx context.Context, filter SMSFilter) ([]SMSContact, error) {
where, args := smsWhere(filter, "m.")
query := `
WITH ranked AS (
WITH resolved AS (
SELECT
m.id, m.device_id, m.imsi, m.peer, m.body, m.message_time,
m.direction,
m.*,
COALESCE(NULLIF(m.modem_imei, ''), 'device:' || m.device_id) AS hardware_key,
COALESCE((
SELECT current_device.id
FROM devices current_device
WHERE m.modem_imei <> ''
AND current_device.modem_imei = m.modem_imei
ORDER BY current_device.updated_at DESC, current_device.id
LIMIT 1
), m.device_id) AS resolved_device_id
FROM sms_messages m` + where + `
), ranked AS (
SELECT
m.id, m.resolved_device_id, m.modem_imei, m.imsi, m.peer,
m.body, m.message_time, m.direction,
ROW_NUMBER() OVER (
PARTITION BY m.device_id, m.imsi, m.peer
PARTITION BY m.hardware_key, m.imsi, m.peer
ORDER BY m.message_time DESC, m.id DESC
) AS row_number,
SUM(CASE
WHEN m.direction IN ('inbound', 'received') AND m.is_read = 0
THEN 1 ELSE 0
END) OVER (
PARTITION BY m.device_id, m.imsi, m.peer
PARTITION BY m.hardware_key, m.imsi, m.peer
) AS unread_count,
COUNT(*) OVER (
PARTITION BY m.device_id, m.imsi, m.peer
PARTITION BY m.hardware_key, m.imsi, m.peer
) AS message_count
FROM sms_messages m` + where + `
FROM resolved m
)
SELECT
r.device_id,
r.resolved_device_id,
COALESCE(d.name, ''),
r.modem_imei,
r.imsi,
COALESCE(NULLIF(dr.phone_number, ''), NULLIF(vr.local_phone, ''), ''),
r.peer,
@@ -482,9 +559,9 @@ func (s *Store) ListSMSContacts(ctx context.Context, filter SMSFilter) ([]SMSCon
r.unread_count,
r.message_count
FROM ranked r
LEFT JOIN devices d ON d.id = r.device_id
LEFT JOIN device_runtime dr ON dr.device_id = r.device_id
LEFT JOIN vowifi_runtime vr ON vr.device_id = r.device_id
LEFT JOIN devices d ON d.id = r.resolved_device_id
LEFT JOIN device_runtime dr ON dr.device_id = r.resolved_device_id
LEFT JOIN vowifi_runtime vr ON vr.device_id = r.resolved_device_id
WHERE r.row_number = 1
ORDER BY r.message_time DESC, r.id DESC
LIMIT ?`
@@ -500,7 +577,7 @@ func (s *Store) ListSMSContacts(ctx context.Context, filter SMSFilter) ([]SMSCon
var value SMSContact
var timestamp int64
if err := rows.Scan(
&value.DeviceID, &value.DeviceName, &value.IMSI,
&value.DeviceID, &value.DeviceName, &value.ModemIMEI, &value.IMSI,
&value.LocalPhone, &value.Peer, &value.DisplayName,
&value.LastMessage, &timestamp, &value.Direction,
&value.LastSMSID, &value.UnreadCount, &value.MessageCount,
@@ -517,7 +594,7 @@ func (s *Store) ListSMSContacts(ctx context.Context, filter SMSFilter) ([]SMSCon
}
const smsMessageSelect = `
SELECT id, message_id, device_id, imsi, peer, direction, body,
SELECT id, message_id, device_id, modem_imei, imsi, peer, direction, body,
message_time, status, source, parts_total, delivery_state, is_read,
extra_json, created_at, updated_at
FROM sms_messages`
@@ -528,7 +605,7 @@ func scanSMSMessage(row rowScanner) (SMSMessage, error) {
var read int
var extra string
err := row.Scan(
&value.ID, &value.MessageID, &value.DeviceID, &value.IMSI,
&value.ID, &value.MessageID, &value.DeviceID, &value.ModemIMEI, &value.IMSI,
&value.Peer, &value.Direction, &value.Body, &messageTime,
&value.Status, &value.Source, &value.PartsTotal,
&value.DeliveryState, &read, &extra, &createdAt, &updatedAt,
@@ -554,6 +631,10 @@ func smsWhere(filter SMSFilter, prefix string) (string, []any) {
clauses = append(clauses, prefix+`device_id = ?`)
args = append(args, filter.DeviceID)
}
if filter.ModemIMEI != "" {
clauses = append(clauses, prefix+`modem_imei = ?`)
args = append(args, filter.ModemIMEI)
}
if filter.IMSI != "" {
clauses = append(clauses, prefix+`imsi = ?`)
args = append(args, filter.IMSI)
@@ -580,6 +661,13 @@ func smsWhere(filter SMSFilter, prefix string) (string, []any) {
return " WHERE " + strings.Join(clauses, " AND "), args
}
func smsHardwareKey(modemIMEI, deviceID string) string {
if modemIMEI = strings.TrimSpace(modemIMEI); modemIMEI != "" {
return modemIMEI
}
return "device:" + strings.TrimSpace(deviceID)
}
func normalizedLimit(value int) int {
if value <= 0 {
return 100
+132
View File
@@ -0,0 +1,132 @@
package store
import (
"encoding/json"
"fmt"
"sort"
"strconv"
"strings"
)
// ConcatMessageIDPrefix marks the stable message id that ingest points assign to
// every segment of one concatenated (long) SMS. Unlike the per-segment modem/IMS
// ids (which embed a storage slot, PDU hash, or RP reference), this id is shared
// by all segments of the message, so SaveSMSMessage folds them into a single row.
const ConcatMessageIDPrefix = "concat:"
// isConcatSMSMessageID reports whether a message id addresses a whole
// concatenated SMS rather than one physical segment.
func isConcatSMSMessageID(messageID string) bool {
return strings.HasPrefix(messageID, ConcatMessageIDPrefix)
}
// ConcatSMSReadyToNotify reports whether an inbound SMS row is ready to surface
// to a notification consumer. A plain message is always ready; a concatenated
// (long) SMS row is ready only once every segment has merged (concat_complete).
// Until then consumers should hold the notification but still advance their
// cursor — the completed message re-enters as a fresh durable id.
func ConcatSMSReadyToNotify(messageID string, extra json.RawMessage) bool {
if !isConcatSMSMessageID(messageID) {
return true
}
document, err := decodeJSONObject(extra)
if err != nil {
return false
}
complete, _ := document["concat_complete"].(bool)
return complete
}
// StableConcatMessageID builds the message id shared by every segment of one
// concatenated SMS. The UDH concat reference is only unique per sender, so the
// hardware identity and peer scope it; total is folded in to further separate the
// rare reference reuse between two different long messages from the same peer.
// The hardware identity matches the row lookup in saveSMSMessage, so a segment
// always finds the row its siblings started.
func StableConcatMessageID(source, modemIMEI, deviceID, peer string, reference, total int) string {
return ConcatMessageIDPrefix + source + ":" + smsHardwareKey(modemIMEI, deviceID) + ":" + peer + ":" +
strconv.Itoa(reference) + ":" + strconv.Itoa(total)
}
// mergeConcatSegment folds one incoming segment into the progressively merged
// body of a concatenated SMS. existingExtra is the stored row's Extra (empty for
// the first segment); segmentBody/segmentExtra are the incoming segment's text
// and Extra, the latter carrying "concat" ({reference,total,sequence}).
//
// Each segment's text is kept under "concat_parts" keyed by its UDH sequence and
// the body is rebuilt by joining the parts in ascending sequence order with no
// separator — exactly how a phone reassembles a long message, and correct for any
// arrival order. The merge is idempotent: redelivering an already-folded sequence
// reports changed=false so callers can skip the write and avoid id churn.
// "concat_complete" flips true once Total segments are present.
func mergeConcatSegment(
existingExtra json.RawMessage,
segmentBody string,
segmentExtra json.RawMessage,
) (body string, extra json.RawMessage, changed bool, err error) {
segment, err := decodeJSONObject(segmentExtra)
if err != nil {
return "", nil, false, fmt.Errorf("decode segment extra: %w", err)
}
concat, _ := segment["concat"].(map[string]any)
sequence := numberAsInt(concat["sequence"])
total := numberAsInt(concat["total"])
if sequence < 1 {
// No usable UDH sequence: keep the incoming segment as the whole body.
return segmentBody, json.RawMessage(segmentExtra), true, nil
}
// Seed the per-segment texts from the previously stored parts so an
// out-of-order arrival always rebuilds in sequence order.
parts := map[int]string{}
if len(existingExtra) > 0 {
if existing, derr := decodeJSONObject(existingExtra); derr == nil {
if stored, ok := existing["concat_parts"].(map[string]any); ok {
for key, value := range stored {
n, aerr := strconv.Atoi(key)
if aerr != nil || n < 1 {
continue
}
if text, ok := value.(string); ok {
parts[n] = text
}
}
}
}
}
prior, alreadyHad := parts[sequence]
changed = !alreadyHad || prior != segmentBody
parts[sequence] = segmentBody
sequences := make([]int, 0, len(parts))
for n := range parts {
sequences = append(sequences, n)
}
sort.Ints(sequences)
var joined strings.Builder
stored := make(map[string]string, len(parts))
for _, n := range sequences {
joined.WriteString(parts[n])
stored[strconv.Itoa(n)] = parts[n]
}
complete := total > 0 && len(parts) >= total
merged := map[string]any{
"concat": concat,
"concat_parts": stored,
"concat_received": len(parts),
"concat_complete": complete,
}
// Preserve non-concat metadata from the latest segment for context.
for _, key := range []string{"encoding", "storage", "transport", "source"} {
if value, ok := segment[key]; ok {
merged[key] = value
}
}
encoded, err := json.Marshal(merged)
if err != nil {
return "", nil, false, fmt.Errorf("encode merged concat extra: %w", err)
}
return joined.String(), json.RawMessage(encoded), changed, nil
}
+229
View File
@@ -0,0 +1,229 @@
package store
import (
"context"
"encoding/json"
"strings"
"testing"
"time"
)
func concatExtra(t *testing.T, reference, total, sequence int) json.RawMessage {
t.Helper()
extra, err := json.Marshal(map[string]any{
"concat": map[string]any{"reference": reference, "total": total, "sequence": sequence},
})
if err != nil {
t.Fatalf("marshal concat extra: %v", err)
}
return extra
}
func TestMergeConcatSegmentJoinsOutOfOrderInSequenceOrder(t *testing.T) {
// UCS-2 long SMS (the customer case) whose segments arrive 2, 1, 3.
var body string
var extra json.RawMessage
var changed, complete bool
body, extra, changed, err := mergeConcatSegment(extra, "中段", concatExtra(t, 9, 3, 2))
if err != nil || !changed {
t.Fatalf("segment 2: body=%q changed=%v err=%v", body, changed, err)
}
if body != "中段" {
t.Fatalf("after segment 2 body = %q, want partial %q", body, "中段")
}
body, extra, changed, err = mergeConcatSegment(extra, "前段", concatExtra(t, 9, 3, 1))
if err != nil || !changed {
t.Fatalf("segment 1: body=%q changed=%v err=%v", body, changed, err)
}
if body != "前段中段" {
t.Fatalf("after segment 1 body = %q, want %q", body, "前段中段")
}
body, extra, changed, err = mergeConcatSegment(extra, "尾段", concatExtra(t, 9, 3, 3))
if err != nil || !changed {
t.Fatalf("segment 3: body=%q changed=%v err=%v", body, changed, err)
}
if body != "前段中段尾段" {
t.Fatalf("complete body = %q, want %q", body, "前段中段尾段")
}
document, err := decodeJSONObject(extra)
if err != nil {
t.Fatalf("decode merged extra: %v", err)
}
complete, _ = document["concat_complete"].(bool)
if !complete {
t.Fatalf("concat_complete = %v, want true; extra=%s", complete, extra)
}
if got := numberAsInt(document["concat_received"]); got != 3 {
t.Fatalf("concat_received = %d, want 3", got)
}
}
func TestMergeConcatSegmentKeepsURLContiguous(t *testing.T) {
// A GSM-7 tracking link split mid-token (the OFCA "garbled" report) must
// reassemble with no break.
first := "https://ofca.gov.hk/track?tok=ab"
second := "cdef1234&lang=zh"
_, extra, _, err := mergeConcatSegment(nil, first, concatExtra(t, 4, 2, 1))
if err != nil {
t.Fatal(err)
}
body, _, _, err := mergeConcatSegment(extra, second, concatExtra(t, 4, 2, 2))
if err != nil {
t.Fatal(err)
}
want := "https://ofca.gov.hk/track?tok=abcdef1234&lang=zh"
if body != want {
t.Fatalf("body = %q, want %q", body, want)
}
}
func TestMergeConcatSegmentRedeliveryIsIdempotent(t *testing.T) {
_, extra, _, err := mergeConcatSegment(nil, "甲", concatExtra(t, 3, 2, 1))
if err != nil {
t.Fatal(err)
}
// A modem rescan redelivers the identical segment: no change, no growth.
body, extra2, changed, err := mergeConcatSegment(extra, "甲", concatExtra(t, 3, 2, 1))
if err != nil {
t.Fatal(err)
}
if changed {
t.Fatalf("redelivered segment reported changed=true; body=%q", body)
}
if body != "甲" {
t.Fatalf("body = %q, want %q", body, "甲")
}
if string(extra2) == "" {
t.Fatal("merged extra lost on idempotent redelivery")
}
}
func TestMergeConcatSegmentWithoutHeaderPassesThrough(t *testing.T) {
extra, err := json.Marshal(map[string]any{"encoding": "gsm7"})
if err != nil {
t.Fatal(err)
}
body, _, changed, err := mergeConcatSegment(nil, "plain", extra)
if err != nil || !changed || body != "plain" {
t.Fatalf("body=%q changed=%v err=%v, want passthrough", body, changed, err)
}
}
func TestStableConcatMessageIDScopesByPeerReferenceTotal(t *testing.T) {
a := StableConcatMessageID("cellular_at", "imei-1", "ec20", "+10086", 7, 2)
if !isConcatSMSMessageID(a) {
t.Fatalf("id %q missing concat prefix", a)
}
if again := StableConcatMessageID("cellular_at", "imei-1", "ec20", "+10086", 7, 2); again != a {
t.Fatalf("id unstable: %q vs %q", a, again)
}
for _, different := range []string{
StableConcatMessageID("cellular_at", "imei-1", "ec20", "+10086", 8, 2), // other reference
StableConcatMessageID("cellular_at", "imei-1", "ec20", "+10010", 7, 2), // other peer
StableConcatMessageID("cellular_at", "imei-1", "ec20", "+10086", 7, 3), // other total
StableConcatMessageID("ims", "imei-1", "ec20", "+10086", 7, 2), // other source
} {
if different == a {
t.Fatalf("id %q collides across distinct concat groups", a)
}
}
}
func TestConcatSMSReadyToNotify(t *testing.T) {
if !ConcatSMSReadyToNotify("modem:SM:3:abcd", json.RawMessage(`{}`)) {
t.Fatal("plain message should always be ready")
}
incomplete := StableConcatMessageID("cellular_at", "imei", "ec20", "peer", 1, 2)
if ConcatSMSReadyToNotify(incomplete, json.RawMessage(`{"concat_complete":false}`)) {
t.Fatal("incomplete long SMS must not notify")
}
if ConcatSMSReadyToNotify(incomplete, json.RawMessage(`not-json`)) {
t.Fatal("unparseable concat extra must not notify")
}
if !ConcatSMSReadyToNotify(incomplete, json.RawMessage(`{"concat_complete":true}`)) {
t.Fatal("complete long SMS should notify")
}
}
func TestSaveConcatSMSFoldsSegmentsIntoOneRow(t *testing.T) {
ctx := context.Background()
database := openTestStore(t, ":memory:")
mustSaveDevice(t, database, "ec20-1", "测试设备")
const imei = "867394042309830"
base := time.Unix(1_700_000_000, 0).UTC()
save := func(sequence int, text string, at time.Time) SMSMessage {
t.Helper()
saved, err := database.SaveSMSMessage(ctx, SMSMessage{
MessageID: StableConcatMessageID("cellular_at", imei, "ec20-1", "+8520000", 5, 2),
DeviceID: "ec20-1", ModemIMEI: imei, IMSI: "45400",
Peer: "+8520000", Direction: "inbound", Body: text,
Timestamp: at, Status: "received", Source: "cellular_at",
PartsTotal: 2,
Extra: concatExtra(t, 5, 2, sequence),
})
if err != nil {
t.Fatalf("SaveSMSMessage(seq=%d) error = %v", sequence, err)
}
return saved
}
first := save(1, "【检测】您的结果为", base)
if first.PartsTotal != 2 {
t.Fatalf("PartsTotal = %d, want 2", first.PartsTotal)
}
if ConcatSMSReadyToNotify(first.MessageID, first.Extra) {
t.Fatal("first segment alone should not be ready to notify")
}
second := save(2, "合格,请查收报告", base.Add(30*time.Second))
if second.ID <= first.ID {
t.Fatalf("completed row id = %d, want a fresh id greater than %d", second.ID, first.ID)
}
if second.Body != "【检测】您的结果为合格,请查收报告" {
t.Fatalf("merged body = %q", second.Body)
}
if !second.Timestamp.Equal(base) {
t.Fatalf("merged timestamp = %v, want earliest segment %v", second.Timestamp, base)
}
if !ConcatSMSReadyToNotify(second.MessageID, second.Extra) {
t.Fatal("completed long SMS should be ready to notify")
}
// Exactly one stored row represents the whole long SMS.
messages, err := database.ListSMSMessages(ctx, SMSFilter{DeviceID: "ec20-1"})
if err != nil {
t.Fatal(err)
}
if len(messages) != 1 {
t.Fatalf("stored rows = %d, want 1 merged row: %+v", len(messages), messages)
}
// The completed row re-enters after the earlier partial id, so the Telegram
// id-cursor surfaces it once, complete.
fresh, err := database.ListInboundSMSAfterID(ctx, first.ID, 10)
if err != nil {
t.Fatal(err)
}
if len(fresh) != 1 || fresh[0].ID != second.ID || !strings.Contains(fresh[0].Body, "合格") {
t.Fatalf("ListInboundSMSAfterID = %+v, want the completed row", fresh)
}
// A modem rescan redelivers an already-folded segment: no write, no id churn.
rescan := save(1, "【检测】您的结果为", base)
if rescan.ID != second.ID {
t.Fatalf("rescan churned the row id: got %d, want stable %d", rescan.ID, second.ID)
}
if rescan.Body != second.Body {
t.Fatalf("rescan changed body to %q", rescan.Body)
}
if after, err := database.ListInboundSMSAfterID(ctx, second.ID, 10); err != nil || len(after) != 0 {
t.Fatalf("rescan produced new rows: %+v, %v", after, err)
}
if count, err := database.ListSMSMessages(ctx, SMSFilter{DeviceID: "ec20-1"}); err != nil || len(count) != 1 {
t.Fatalf("rows after rescan = %d, %v; want still 1", len(count), err)
}
}
+9 -1
View File
@@ -13,7 +13,7 @@ import (
_ "modernc.org/sqlite"
)
const schemaVersion = 6
const schemaVersion = 8
var ErrNotFound = errors.New("store: not found")
@@ -117,6 +117,14 @@ func migrate(ctx context.Context, db *sql.DB) error {
}
for _, statement := range migrationStatements(nextVersion) {
if _, err := tx.ExecContext(ctx, statement); err != nil {
// A database whose user_version was repaired or rolled back may
// already contain an additive column. Remaining statements in the
// migration are still safe and must be applied.
duplicateAdditiveColumn := (nextVersion == 7 && strings.Contains(statement, "ADD COLUMN modem_imei")) ||
(nextVersion == 8 && strings.Contains(statement, "ADD COLUMN device_type"))
if duplicateAdditiveColumn && strings.Contains(strings.ToLower(err.Error()), "duplicate column name") {
continue
}
_ = tx.Rollback()
return fmt.Errorf("apply sqlite migration %d: %w", nextVersion, err)
}
+1 -1
View File
@@ -13,7 +13,7 @@ func TestAssetNamesFor(t *testing.T) {
}{
{"linux", "amd64", []string{"vocat-linux-amd64"}},
{"linux", "386", []string{"vocat-linux-386"}},
{"linux", "arm64", []string{"vocat-linux-arm64"}},
{"linux", "arm64", []string{"vocat-linux-arm64", "vocat-linux-aarch64"}},
{"linux", "arm", []string{"vocat-linux-armv7", "vocat-linux-arm"}},
}
for _, item := range tests {
+35 -1
View File
@@ -25,7 +25,20 @@ type Asset struct {
Size int64 `json:"size"`
}
const githubAPI = "https://api.github.com"
// CheckResult describes a trusted release check without downloading assets.
type CheckResult struct {
Available bool
Applied bool
Current string
Latest string
ReleaseNotes string
Release *Release
}
const (
githubAPI = "https://api.github.com"
DefaultRepository = "MengMengCode/VoCat"
)
// LatestRelease fetches the newest published release for repo (form
// "owner/name"). A non-empty token is sent as a Bearer header, which is
@@ -73,6 +86,27 @@ func LatestRelease(ctx context.Context, repo, token string) (*Release, error) {
return &release, nil
}
// CheckLatest fetches the newest release and performs a semantic version
// comparison so development builds are never offered an older release.
func CheckLatest(ctx context.Context, repo, token, current string) (CheckResult, error) {
release, err := LatestRelease(ctx, repo, token)
if err != nil {
return CheckResult{}, err
}
latest := strings.TrimPrefix(strings.TrimSpace(release.TagName), "v")
available, err := IsNewerVersion(current, latest)
if err != nil {
return CheckResult{}, fmt.Errorf("update: compare release versions: %w", err)
}
return CheckResult{
Available: available,
Current: current,
Latest: latest,
ReleaseNotes: strings.TrimSpace(release.Body),
Release: release,
}, nil
}
// downloadAsset streams a release asset into dst, honoring the request context.
// The token is applied for consistency with the API call (GitHub release assets
// redirect to a pre-signed S3 URL; the token is dropped on redirect, which is
+59 -27
View File
@@ -7,8 +7,8 @@
// Trust model: GitHub TLS guarantees the channel; the repository owner controls
// which assets are published; SHA256SUMS guards integrity. There is no GPG
// signature verification — an accepted trade-off for a closed-network testing
// tool. The web UI's check-update button remains an intentional no-op; only the
// CLI performs code replacement.
// tool. Both the CLI and authenticated web UI use this same verified replacement
// path.
package update
import (
@@ -51,12 +51,12 @@ func Run(logger *slog.Logger, args []string) error {
if opts.Repo == "" {
opts.Repo = strings.TrimSpace(os.Getenv("VOCAT_REPO"))
}
if opts.Repo == "" {
opts.Repo = DefaultRepository
}
if opts.Token == "" {
opts.Token = strings.TrimSpace(os.Getenv("GITHUB_TOKEN"))
}
if opts.Repo == "" {
return fmt.Errorf("update: no repository configured (set --repo=owner/name or VOCAT_REPO)")
}
if opts.Target == "" {
opts.Target = resolveDefaultTarget()
}
@@ -65,33 +65,55 @@ func Run(logger *slog.Logger, args []string) error {
defer cancel()
logger.Info("checking for updates", "repo", opts.Repo, "current", buildinfo.Version)
release, err := LatestRelease(ctx, opts.Repo, opts.Token)
result, err := CheckLatest(ctx, opts.Repo, opts.Token, buildinfo.Version)
if err != nil {
return err
}
latest := strings.TrimPrefix(release.TagName, "v")
if latest == "" {
latest = release.TagName
}
if latest == buildinfo.Version && !opts.Force {
if !result.Available && !opts.Force {
logger.Info("already up to date", "version", buildinfo.Version)
fmt.Printf("vocat %s is already the latest release.\n", buildinfo.Version)
return nil
}
if opts.Check {
fmt.Printf("update available: %s -> %s\n", buildinfo.Version, latest)
if release.Body != "" {
fmt.Println(strings.TrimSpace(release.Body))
fmt.Printf("update available: %s -> %s\n", buildinfo.Version, result.Latest)
if result.ReleaseNotes != "" {
fmt.Println(result.ReleaseNotes)
}
return nil
}
logger.Info("update available", "current", buildinfo.Version, "latest", latest)
return applyUpdate(ctx, logger, opts, release, latest)
logger.Info("update available", "current", buildinfo.Version, "latest", result.Latest)
return applyUpdate(ctx, logger, opts, result.Release, result.Latest, true)
}
func applyUpdate(ctx context.Context, logger *slog.Logger, opts Options, release *Release, latest string) error {
// ApplyLatest downloads, verifies, and atomically installs the newest trusted
// release. HTTP callers can pass restart=false and restart after flushing the
// response.
func ApplyLatest(ctx context.Context, logger *slog.Logger, opts Options, restart bool) (CheckResult, error) {
if strings.TrimSpace(opts.Repo) == "" {
opts.Repo = DefaultRepository
}
if strings.TrimSpace(opts.Token) == "" {
opts.Token = strings.TrimSpace(os.Getenv("GITHUB_TOKEN"))
}
if strings.TrimSpace(opts.Target) == "" {
opts.Target = resolveDefaultTarget()
}
result, err := CheckLatest(ctx, opts.Repo, opts.Token, buildinfo.Version)
if err != nil {
return CheckResult{}, err
}
if !result.Available && !opts.Force {
return result, nil
}
if err := applyUpdate(ctx, logger, opts, result.Release, result.Latest, restart); err != nil {
return CheckResult{}, err
}
result.Applied = true
return result, nil
}
func applyUpdate(ctx context.Context, logger *slog.Logger, opts Options, release *Release, latest string, restart bool) error {
assetNames := assetNamesFor(runtime.GOOS, runtime.GOARCH)
var asset *Asset
for _, name := range assetNames {
@@ -169,11 +191,13 @@ func applyUpdate(ctx context.Context, logger *slog.Logger, opts Options, release
logger.Info("installed new binary", "target", opts.Target, "version", latest)
fmt.Printf("vocat updated to %s.\n", latest)
if err := restartService(logger); err != nil {
// The file replacement already succeeded; a restart failure is not
// fatal — the operator can restart the service manually.
fmt.Printf("Binary replaced, but automatic restart failed: %v\n", err)
fmt.Println("Restart the vocat service manually to apply the new build.")
if restart {
if err := RestartService(logger); err != nil {
// The file replacement already succeeded; a restart failure is not
// fatal — the operator can restart the service manually.
fmt.Printf("Binary replaced, but automatic restart failed: %v\n", err)
fmt.Println("Restart the vocat service manually to apply the new build.")
}
}
return nil
}
@@ -201,14 +225,17 @@ func backupAndReplace(target, tmp string) error {
return nil
}
// restartService restarts the vocat systemd unit. If systemctl is unavailable
// RestartService restarts the vocat systemd unit. If systemctl is unavailable
// (non-systemd hosts, containers), it returns an error the caller surfaces as
// a non-fatal warning.
func restartService(logger *slog.Logger) error {
func RestartService(logger *slog.Logger) error {
if _, err := exec.LookPath("systemctl"); err != nil {
return fmt.Errorf("systemctl not found in PATH")
}
cmd := exec.Command("systemctl", "restart", "vocat")
// Queue the restart and let systemctl exit before systemd stops this unit.
// A blocking restart command becomes part of vocat.service's own cgroup and
// waits for that same cgroup to terminate, creating a stop-timeout cycle.
cmd := exec.Command("systemctl", "restart", "--no-block", "vocat")
if out, err := cmd.CombinedOutput(); err != nil {
logger.Warn("systemctl restart failed", "error", err, "output", string(out))
return fmt.Errorf("systemctl restart vocat: %w", err)
@@ -245,6 +272,11 @@ func findAsset(release *Release, name string) *Asset {
}
func assetNamesFor(goos, goarch string) []string {
if goos == "linux" && goarch == "arm64" {
// AArch64 and arm64 name the same instruction set. Prefer the historic
// release name and accept the explicit architecture alias as fallback.
return []string{"vocat-linux-arm64", "vocat-linux-aarch64"}
}
if goos == "linux" && goarch == "arm" {
// Official 32-bit ARM builds target GOARM=7. Keep the generic legacy
// name as a fallback for installations consuming an older release.
@@ -261,7 +293,7 @@ Fetch the latest release from GitHub and replace this binary in place.
Flags:
--check Report whether an update is available, then exit.
--force Reinstall even when already at the latest version.
--repo owner/name GitHub repository (default: $VOCAT_REPO).
--repo owner/name GitHub repository (default: $VOCAT_REPO or MengMengCode/VoCat).
--target path Binary to replace (default: /opt/vocat/bin/vocat if
present, otherwise the running executable).
--token token GitHub bearer token (default: $GITHUB_TOKEN).
+120
View File
@@ -0,0 +1,120 @@
package update
import (
"fmt"
"strconv"
"strings"
)
type semanticVersion struct {
major int
minor int
patch int
prerelease string
}
// IsNewerVersion reports whether latest is newer than current. Both values
// may include the conventional v prefix, prerelease suffixes, and build
// metadata. Invalid release versions are rejected instead of triggering a
// downgrade or an arbitrary file replacement.
func IsNewerVersion(current, latest string) (bool, error) {
currentVersion, err := parseSemanticVersion(current)
if err != nil {
return false, fmt.Errorf("current version: %w", err)
}
latestVersion, err := parseSemanticVersion(latest)
if err != nil {
return false, fmt.Errorf("latest version: %w", err)
}
if currentVersion.major != latestVersion.major {
return latestVersion.major > currentVersion.major, nil
}
if currentVersion.minor != latestVersion.minor {
return latestVersion.minor > currentVersion.minor, nil
}
if currentVersion.patch != latestVersion.patch {
return latestVersion.patch > currentVersion.patch, nil
}
if currentVersion.prerelease == latestVersion.prerelease {
return false, nil
}
if currentVersion.prerelease != "" && latestVersion.prerelease == "" {
return true, nil
}
if currentVersion.prerelease == "" {
return false, nil
}
return comparePrerelease(currentVersion.prerelease, latestVersion.prerelease) < 0, nil
}
func parseSemanticVersion(raw string) (semanticVersion, error) {
value := strings.TrimPrefix(strings.TrimSpace(raw), "v")
if build := strings.IndexByte(value, '+'); build >= 0 {
value = value[:build]
}
prerelease := ""
hasPrerelease := false
if dash := strings.IndexByte(value, '-'); dash >= 0 {
hasPrerelease = true
prerelease = value[dash+1:]
value = value[:dash]
}
parts := strings.Split(value, ".")
if len(parts) != 3 || (hasPrerelease && prerelease == "") {
return semanticVersion{}, fmt.Errorf("%q is not a semantic version", raw)
}
numbers := make([]int, 3)
for index, part := range parts {
if part == "" || (len(part) > 1 && part[0] == '0') {
return semanticVersion{}, fmt.Errorf("%q is not a semantic version", raw)
}
value, err := strconv.Atoi(part)
if err != nil || value < 0 {
return semanticVersion{}, fmt.Errorf("%q is not a semantic version", raw)
}
numbers[index] = value
}
if strings.ContainsAny(prerelease, " \t\r\n") {
return semanticVersion{}, fmt.Errorf("%q is not a semantic version", raw)
}
return semanticVersion{
major: numbers[0],
minor: numbers[1],
patch: numbers[2],
prerelease: prerelease,
}, nil
}
func comparePrerelease(left, right string) int {
leftParts := strings.Split(left, ".")
rightParts := strings.Split(right, ".")
for index := 0; index < len(leftParts) && index < len(rightParts); index++ {
if leftParts[index] == rightParts[index] {
continue
}
leftNumber, leftErr := strconv.Atoi(leftParts[index])
rightNumber, rightErr := strconv.Atoi(rightParts[index])
switch {
case leftErr == nil && rightErr == nil:
if leftNumber < rightNumber {
return -1
}
return 1
case leftErr == nil:
return -1
case rightErr == nil:
return 1
case leftParts[index] < rightParts[index]:
return -1
default:
return 1
}
}
if len(leftParts) < len(rightParts) {
return -1
}
if len(leftParts) > len(rightParts) {
return 1
}
return 0
}
+34
View File
@@ -0,0 +1,34 @@
package update
import "testing"
func TestIsNewerVersion(t *testing.T) {
tests := []struct {
current string
latest string
want bool
}{
{"0.0.3", "v0.0.4", true},
{"0.0.4", "v0.0.4", false},
{"0.1.0-dev", "v0.0.4", false},
{"0.1.0-dev", "v0.1.0", true},
{"1.2.3-rc.1", "v1.2.3-rc.2", true},
{"1.2.3", "v1.2.3-rc.2", false},
}
for _, item := range tests {
got, err := IsNewerVersion(item.current, item.latest)
if err != nil {
t.Errorf("IsNewerVersion(%q, %q): %v", item.current, item.latest, err)
continue
}
if got != item.want {
t.Errorf("IsNewerVersion(%q, %q) = %v, want %v", item.current, item.latest, got, item.want)
}
}
}
func TestIsNewerVersionRejectsInvalidRelease(t *testing.T) {
if _, err := IsNewerVersion("0.1.0", "nightly"); err == nil {
t.Fatal("invalid latest version was accepted")
}
}
+10 -1
View File
@@ -87,6 +87,13 @@ func (relay *sessionRelay) run() {
packet := append([]byte(nil), buffer[:n]...)
if isIKE {
if err := relay.handleIKE(packet); err != nil {
if errors.Is(err, errMismatchedSessionSPIs) {
// A reconnect can reuse the same NAT mapping while the ePDG still
// has packets queued for the previous IKE SA. Those packets are
// unrelated to this authenticated session and must be discarded;
// treating one as fatal tears down the newly established CHILD_SA.
continue
}
relay.fail(err)
return
}
@@ -111,13 +118,15 @@ func (relay *sessionRelay) run() {
}
}
var errMismatchedSessionSPIs = errors.New("ike: session packet has mismatched SPIs")
func (relay *sessionRelay) handleIKE(packet []byte) error {
header, _, err := parseIKEPacket(packet)
if err != nil {
return err
}
if header.InitiatorSPI != relay.spii || header.ResponderSPI != relay.spir {
return errors.New("ike: session packet has mismatched SPIs")
return errMismatchedSessionSPIs
}
if header.Flags&flagResponse != 0 {
return nil
+36
View File
@@ -178,3 +178,39 @@ func TestSessionRelaySendsNATKeepalive(t *testing.T) {
t.Fatal("relay did not send a NAT-T keepalive")
}
}
func TestSessionRelayDropsDelayedIKEPacketFromPreviousSA(t *testing.T) {
transport := newFakeSessionTransport()
spii := [8]byte{1}
spir := [8]byte{2}
relay := newSessionRelay(
transport,
legacyTestSuite(),
ikeKeys{},
spii,
spir,
true,
time.Hour,
)
defer relay.Close()
transport.incoming <- fakeSessionPacket{
ike: true,
data: ikeHeader{
InitiatorSPI: [8]byte{9},
ResponderSPI: [8]byte{8},
Exchange: exchangeInformational,
}.marshal(nil),
}
wantedESP := []byte{0, 0, 0, 9, 0, 0, 0, 1, 0xaa}
transport.incoming <- fakeSessionPacket{data: wantedESP}
buffer := make([]byte, 64)
count, err := relay.ReceiveESP(context.Background(), buffer)
if err != nil {
t.Fatalf("ReceiveESP() after stale IKE packet = %v", err)
}
if !bytes.Equal(buffer[:count], wantedESP) {
t.Fatalf("ESP after stale IKE packet = %x, want %x", buffer[:count], wantedESP)
}
}
+569
View File
@@ -0,0 +1,569 @@
package ims
import (
"context"
"errors"
"fmt"
"net"
"sort"
"strconv"
"strings"
"time"
"vocat/internal/vowifi"
)
var (
ErrCallNotFound = errors.New("ims: call not found")
ErrCallState = errors.New("ims: call is not in the required state")
)
const terminalCallRetention = 30 * time.Second
type imsCall struct {
public vowifi.Call
callID string
target string
from string
to string
branch string
cseq uint32
invite *sipRequest
respond func([]byte) error
responses chan *sipResponse
remoteTag string
routes []string
terminated bool
media *rtpMedia
}
func (session *Session) Calls() []vowifi.Call {
session.callMu.Lock()
defer session.callMu.Unlock()
now := time.Now().UTC()
calls := make([]vowifi.Call, 0, len(session.calls))
for id, call := range session.calls {
if call.public.EndedAt != nil && now.Sub(*call.public.EndedAt) > terminalCallRetention {
delete(session.calls, id)
continue
}
calls = append(calls, call.public)
}
sort.Slice(calls, func(i, j int) bool { return calls[i].StartedAt.Before(calls[j].StartedAt) })
return calls
}
func (session *Session) DialCall(ctx context.Context, number string) (vowifi.Call, error) {
number = strings.TrimSpace(number)
if !validCallNumber(number) {
return vowifi.Call{}, errors.New("ims: invalid dial number")
}
callToken, err := randomHex(18)
if err != nil {
return vowifi.Call{}, err
}
branch, err := randomHex(12)
if err != nil {
return vowifi.Call{}, err
}
callID := callToken + "@" + addressHost(session.conn.LocalAddr())
target := "tel:" + number
session.mu.Lock()
cseq := session.cseq
session.cseq++
routes := append([]string(nil), session.evidence.ServiceRoute...)
securityHeaders := runtimeSecurityHeaders(session.securityActive, session.securityAgreement.verifyValue)
session.mu.Unlock()
media, err := newRTPMedia(session.localMediaIP())
if err != nil {
return vowifi.Call{}, err
}
body := media.offerSDP(session.localMediaIP())
transportUpper := strings.ToUpper(session.transport)
from := "<" + session.identity.public + ">;tag=" + session.fromTag
to := "<" + target + ">"
lines := []string{
"INVITE " + target + " SIP/2.0",
fmt.Sprintf("Via: SIP/2.0/%s %s;branch=z9hG4bK%s;rport", transportUpper, session.conn.LocalAddr().String(), branch),
"Max-Forwards: 70",
}
lines = append(lines, securityHeaders...)
if len(routes) == 0 {
lines = append(lines, "Route: <sip:"+session.endpoint.address()+";transport="+session.transport+";lr>")
} else {
for _, route := range routes {
lines = append(lines, "Route: "+route)
}
}
lines = append(lines,
"From: "+from,
"To: "+to,
"Call-ID: "+callID,
fmt.Sprintf("CSeq: %d INVITE", cseq),
"Contact: <sip:"+session.identity.user+"@"+session.contactAddress()+";transport="+session.transport+">",
"P-Preferred-Identity: <"+session.identity.public+">",
"Allow: INVITE, ACK, CANCEL, BYE, OPTIONS, MESSAGE",
"Supported: timer",
"Content-Type: application/sdp",
"Content-Length: "+strconv.Itoa(len(body)), "", "",
)
request := append([]byte(strings.Join(lines, "\r\n")), body...)
responses := make(chan *sipResponse, 8)
key := sipTransactionKey{callID: callID, cseq: cseq, method: "INVITE"}
session.transactionsMu.Lock()
if _, duplicate := session.transactions[key]; duplicate {
session.transactionsMu.Unlock()
_ = media.Close()
return vowifi.Call{}, errors.New("ims: duplicate call transaction")
}
session.transactions[key] = responses
session.transactionsMu.Unlock()
call := &imsCall{
public: vowifi.Call{ID: callID, Number: number, Direction: "outgoing", State: "dialing", StartedAt: time.Now().UTC()},
callID: callID, target: target, from: from, to: to, branch: branch, cseq: cseq, responses: responses,
routes: routes, media: media,
}
session.callMu.Lock()
session.calls[callID] = call
session.callMu.Unlock()
session.writeMu.Lock()
_, err = session.conn.Write(request)
session.writeMu.Unlock()
if err != nil {
_ = media.Close()
session.transactionsMu.Lock()
delete(session.transactions, key)
session.transactionsMu.Unlock()
session.callMu.Lock()
delete(session.calls, callID)
session.callMu.Unlock()
return vowifi.Call{}, fmt.Errorf("ims: send SIP INVITE: %w", err)
}
go session.watchOutgoingCall(call, key)
return call.public, nil
}
func (session *Session) watchOutgoingCall(call *imsCall, key sipTransactionKey) {
timer := time.NewTimer(2 * time.Minute)
defer timer.Stop()
defer func() {
session.transactionsMu.Lock()
delete(session.transactions, key)
session.transactionsMu.Unlock()
}()
for {
select {
case <-session.refreshContext.Done():
return
case <-timer.C:
session.finishCall(call.callID, "failed", 0, "SIP INVITE transaction timed out")
return
case response := <-call.responses:
if response == nil {
continue
}
if response.StatusCode < 200 {
session.setCallDiagnostic(call.callID, response.StatusCode, response.Reason)
if response.StatusCode >= 180 {
session.setCallState(call.callID, "ringing")
}
continue
}
if response.StatusCode >= 200 && response.StatusCode < 300 {
session.callMu.Lock()
call.to = response.value("To")
call.remoteTag = headerParameter(call.to, "tag")
if contact := headerURI(response.value("Contact")); contact != "" {
call.target = contact
}
if recordRoutes := response.values("Record-Route"); len(recordRoutes) > 0 {
call.routes = reverseStrings(recordRoutes)
}
session.callMu.Unlock()
mediaErr := call.media.configureRemote(response.Body)
_ = session.sendACK(call)
if mediaErr != nil {
session.finishCall(call.callID, "failed", response.StatusCode, mediaErr.Error())
go func() {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_ = session.sendDialogRequest(ctx, call, "BYE")
}()
return
}
session.setCallMediaReady(call.callID)
session.setCallState(call.callID, "active")
} else {
session.finishCall(call.callID, "failed", response.StatusCode, response.Reason)
}
return
}
}
}
func (session *Session) AnswerCall(_ context.Context, id string) (vowifi.Call, error) {
session.callMu.Lock()
call := session.calls[id]
if call == nil {
session.callMu.Unlock()
return vowifi.Call{}, ErrCallNotFound
}
if call.public.Direction != "incoming" || call.public.State != "ringing" || call.invite == nil || call.respond == nil {
session.callMu.Unlock()
return vowifi.Call{}, ErrCallState
}
request, respond := call.invite, call.respond
session.callMu.Unlock()
response, err := buildSIPResponseWithBody(request, 200, session.fromTag, call.media.answerSDP(session.localMediaIP()))
if err != nil {
return vowifi.Call{}, err
}
if err := respond(response); err != nil {
return vowifi.Call{}, err
}
session.setCallState(id, "active")
if call.media.ready() {
session.setCallMediaReady(id)
}
session.callMu.Lock()
result := call.public
session.callMu.Unlock()
return result, nil
}
func (session *Session) HangupCall(ctx context.Context, id string) error {
session.callMu.Lock()
call := session.calls[id]
if call == nil {
session.callMu.Unlock()
return ErrCallNotFound
}
state := call.public.State
direction := call.public.Direction
request, respond := call.invite, call.respond
session.callMu.Unlock()
if direction == "incoming" && state == "ringing" && request != nil && respond != nil {
response, err := buildSIPResponseWithBody(request, 486, session.fromTag, nil)
if err != nil {
return err
}
if err := respond(response); err != nil {
return err
}
session.finishCall(id, "ended", 0, "")
return nil
}
method := "BYE"
if direction == "outgoing" && (state == "dialing" || state == "ringing") {
method = "CANCEL"
}
err := session.sendDialogRequest(ctx, call, method)
// A remote endpoint may already have removed the dialog and answer BYE with
// 481. The local call must still leave the active list after a hang-up.
session.finishCall(id, "ended", 0, "")
return err
}
func (session *Session) handleCallRequest(request *sipRequest, respond func([]byte) error) bool {
switch request.Method {
case "INVITE":
callID := strings.TrimSpace(request.value("Call-ID"))
if callID == "" {
return true
}
number := identityNumber(request.value("From"))
target := headerURI(request.value("Contact"))
if target == "" {
target = request.URI
}
media, err := newRTPMedia(session.localMediaIP())
if err != nil {
if response, buildErr := buildSIPResponseWithBody(request, 488, session.fromTag, nil); buildErr == nil {
_ = respond(response)
}
return true
}
if len(request.Body) > 0 {
if err := media.configureRemote(request.Body); err != nil {
_ = media.Close()
if response, buildErr := buildSIPResponseWithBody(request, 488, session.fromTag, nil); buildErr == nil {
_ = respond(response)
}
return true
}
}
call := &imsCall{
public: vowifi.Call{ID: callID, Number: number, Direction: "incoming", State: "ringing", StartedAt: time.Now().UTC()},
callID: callID, target: target, from: request.value("To") + ";tag=" + session.fromTag,
to: request.value("From"), invite: request, respond: respond, routes: request.values("Record-Route"), media: media,
}
session.callMu.Lock()
session.calls[callID] = call
session.callMu.Unlock()
response, err := buildSIPResponseWithBody(request, 180, session.fromTag, nil)
if err == nil {
_ = respond(response)
}
return true
case "ACK":
callID := strings.TrimSpace(request.value("Call-ID"))
session.callMu.Lock()
call := session.calls[callID]
session.callMu.Unlock()
if call != nil && call.media != nil && !call.media.ready() && len(request.Body) > 0 {
if err := call.media.configureRemote(request.Body); err != nil {
session.finishCall(callID, "failed", 0, err.Error())
} else {
session.setCallMediaReady(callID)
}
}
return true
case "CANCEL", "BYE":
response, err := buildSIPResponseWithBody(request, 200, session.fromTag, nil)
if err == nil {
_ = respond(response)
}
callID := strings.TrimSpace(request.value("Call-ID"))
if request.Method == "CANCEL" {
session.callMu.Lock()
call := session.calls[callID]
session.callMu.Unlock()
if call != nil && call.invite != nil && call.respond != nil {
if terminated, buildErr := buildSIPResponseWithBody(call.invite, 487, session.fromTag, nil); buildErr == nil {
_ = call.respond(terminated)
}
}
}
session.finishCall(callID, "ended", 0, "")
return true
default:
return false
}
}
func (session *Session) sendACK(call *imsCall) error {
request := session.buildDialogRequest(call, "ACK", call.cseq)
session.writeMu.Lock()
_, err := session.conn.Write(request)
session.writeMu.Unlock()
return err
}
func (session *Session) sendDialogRequest(ctx context.Context, call *imsCall, method string) error {
cseq := call.cseq
if method == "BYE" {
session.mu.Lock()
cseq = session.cseq
session.cseq++
session.mu.Unlock()
}
request := session.buildDialogRequest(call, method, cseq)
if method == "ACK" {
session.writeMu.Lock()
_, err := session.conn.Write(request)
session.writeMu.Unlock()
return err
}
response, err := session.exchangeRuntime(ctx, request, sipTransactionKey{callID: call.callID, cseq: cseq, method: method})
if err != nil {
return err
}
if response.StatusCode < 200 || response.StatusCode >= 300 {
return fmt.Errorf("ims: SIP %s rejected with %d", method, response.StatusCode)
}
return nil
}
func (session *Session) buildDialogRequest(call *imsCall, method string, cseq uint32) []byte {
branch, _ := randomHex(12)
if method == "CANCEL" {
branch = call.branch
}
to := call.to
if to == "" {
to = "<" + call.target + ">"
}
lines := []string{
method + " " + call.target + " SIP/2.0",
fmt.Sprintf("Via: SIP/2.0/%s %s;branch=z9hG4bK%s;rport", strings.ToUpper(session.transport), session.conn.LocalAddr().String(), branch),
"Max-Forwards: 70",
}
for _, route := range call.routes {
lines = append(lines, "Route: "+route)
}
lines = append(lines,
"From: "+call.from,
"To: "+to,
"Call-ID: "+call.callID,
fmt.Sprintf("CSeq: %d %s", cseq, method),
"Content-Length: 0", "", "",
)
return []byte(strings.Join(lines, "\r\n"))
}
func (session *Session) localMediaIP() net.IP {
var localAddress net.Addr
if session.conn != nil {
localAddress = session.conn.LocalAddr()
}
return addressIP(localAddress)
}
func buildSIPResponseWithBody(request *sipRequest, status int, tag string, body []byte) ([]byte, error) {
reasons := map[int]string{180: "Ringing", 200: "OK", 486: "Busy Here", 487: "Request Terminated", 488: "Not Acceptable Here"}
reason := reasons[status]
if reason == "" {
return nil, errors.New("ims: unsupported call response status")
}
via := request.values("Via")
from, to := request.value("From"), request.value("To")
callID, cseq := request.value("Call-ID"), request.value("CSeq")
if len(via) == 0 || from == "" || to == "" || callID == "" || cseq == "" {
return nil, errors.New("ims: call request omitted a mandatory response header")
}
if !strings.Contains(strings.ToLower(to), ";tag=") {
to += ";tag=" + tag
}
lines := []string{fmt.Sprintf("SIP/2.0 %d %s", status, reason)}
for _, value := range via {
lines = append(lines, "Via: "+value)
}
lines = append(lines, "From: "+from, "To: "+to, "Call-ID: "+callID, "CSeq: "+cseq)
if len(body) > 0 {
lines = append(lines, "Content-Type: application/sdp")
}
lines = append(lines, "Content-Length: "+strconv.Itoa(len(body)), "", "")
return append([]byte(strings.Join(lines, "\r\n")), body...), nil
}
func (session *Session) setCallState(id, state string) {
session.callMu.Lock()
if call := session.calls[id]; call != nil {
call.public.State = state
if state != "ended" && state != "failed" {
call.public.EndedAt = nil
}
}
session.callMu.Unlock()
}
func (session *Session) setCallDiagnostic(id string, code int, reason string) {
session.callMu.Lock()
if call := session.calls[id]; call != nil {
call.public.SIPCode = code
call.public.Reason = safeSIPDiagnostic(reason)
}
session.callMu.Unlock()
}
func (session *Session) setCallMediaReady(id string) {
session.callMu.Lock()
if call := session.calls[id]; call != nil && call.media != nil {
call.public.MediaReady = call.media.ready()
call.public.Codec = call.media.Codec()
}
session.callMu.Unlock()
}
func (session *Session) CallMedia(_ context.Context, id string) (vowifi.CallMedia, error) {
session.callMu.Lock()
defer session.callMu.Unlock()
call := session.calls[id]
if call == nil {
return nil, ErrCallNotFound
}
if call.public.State != "active" || call.media == nil || !call.media.ready() {
return nil, ErrCallState
}
return call.media, nil
}
func (session *Session) finishCall(id, state string, code int, reason string) {
now := time.Now().UTC()
var media *rtpMedia
session.callMu.Lock()
if call := session.calls[id]; call != nil {
media = call.media
call.public.State = state
if code != 0 {
call.public.SIPCode = code
}
if reason = safeSIPDiagnostic(reason); reason != "" {
call.public.Reason = reason
}
call.public.EndedAt = &now
}
session.callMu.Unlock()
if media != nil {
_ = media.Close()
}
}
func validCallNumber(value string) bool {
if len(value) < 2 || len(value) > 32 {
return false
}
for index, character := range value {
if character >= '0' && character <= '9' || index == 0 && character == '+' || character == '*' || character == '#' {
continue
}
return false
}
return true
}
func identityNumber(value string) string {
value = strings.TrimSpace(value)
if start := strings.Index(value, "<"); start >= 0 {
if end := strings.Index(value[start:], ">"); end > 0 {
value = value[start+1 : start+end]
}
}
value = strings.TrimPrefix(value, "sip:")
value = strings.TrimPrefix(value, "tel:")
if at := strings.Index(value, "@"); at >= 0 {
value = value[:at]
}
return strings.TrimSpace(value)
}
func headerParameter(value, name string) string {
needle := ";" + strings.ToLower(name) + "="
lower := strings.ToLower(value)
index := strings.Index(lower, needle)
if index < 0 {
return ""
}
value = value[index+len(needle):]
if end := strings.IndexAny(value, ";,> \t"); end >= 0 {
value = value[:end]
}
return strings.Trim(value, `"`)
}
func headerURI(value string) string {
value = strings.TrimSpace(value)
if start := strings.Index(value, "<"); start >= 0 {
if end := strings.Index(value[start+1:], ">"); end >= 0 {
return strings.TrimSpace(value[start+1 : start+1+end])
}
}
if end := strings.Index(value, ";"); end >= 0 {
value = value[:end]
}
if strings.HasPrefix(strings.ToLower(value), "sip:") || strings.HasPrefix(strings.ToLower(value), "tel:") {
return strings.TrimSpace(value)
}
return ""
}
func reverseStrings(values []string) []string {
result := append([]string(nil), values...)
for left, right := 0, len(result)-1; left < right; left, right = left+1, right-1 {
result[left], result[right] = result[right], result[left]
}
return result
}
var _ vowifi.CallController = (*Session)(nil)
var _ vowifi.CallMediaController = (*Session)(nil)
+87
View File
@@ -0,0 +1,87 @@
package ims
import (
"context"
"strings"
"testing"
"vocat/internal/vowifi"
)
func TestIncomingCallCanRingAndAnswerWithMediaOffer(t *testing.T) {
session := &Session{fromTag: "local-tag", calls: make(map[string]*imsCall)}
packet, err := parseSIPPacket([]byte(strings.Join([]string{
"INVITE sip:[email protected] SIP/2.0",
"Via: SIP/2.0/UDP 192.0.2.10:5060;branch=z9hG4bK-incoming",
"From: <tel:+447700900001>;tag=remote",
"To: <sip:[email protected]>",
"Call-ID: [email protected]",
"CSeq: 1 INVITE",
"Content-Length: 0", "", "",
}, "\r\n")))
if err != nil || packet.Request == nil {
t.Fatalf("parse INVITE: %v", err)
}
var responses [][]byte
session.handleSIPRequest(packet.Request, func(response []byte) error {
responses = append(responses, append([]byte(nil), response...))
return nil
})
calls := session.Calls()
if len(calls) != 1 || calls[0].Direction != "incoming" || calls[0].State != "ringing" || calls[0].Number != "+447700900001" {
t.Fatalf("incoming Calls = %#v", calls)
}
if len(responses) != 1 || !strings.HasPrefix(string(responses[0]), "SIP/2.0 180 Ringing") {
t.Fatalf("ringing response = %q", responses)
}
answered, err := session.AnswerCall(context.Background(), calls[0].ID)
if err != nil {
t.Fatal(err)
}
if answered.State != "active" || len(responses) != 2 || !strings.Contains(string(responses[1]), "a=sendrecv") {
t.Fatalf("answered = %#v, response = %q", answered, responses[1])
}
}
func TestIncomingCallCanBeRejected(t *testing.T) {
session := &Session{fromTag: "local-tag", calls: make(map[string]*imsCall)}
packet, err := parseSIPPacket([]byte(strings.Join([]string{
"INVITE sip:[email protected] SIP/2.0",
"Via: SIP/2.0/UDP 192.0.2.10:5060;branch=z9hG4bK-a",
"From: <tel:+1>;tag=a", "To: <sip:[email protected]>",
"Call-ID: reject-call", "CSeq: 1 INVITE", "Content-Length: 0", "", "",
}, "\r\n")))
if err != nil || packet.Request == nil {
t.Fatalf("parse INVITE: %v", err)
}
var response []byte
session.handleCallRequest(packet.Request, func(value []byte) error { response = append([]byte(nil), value...); return nil })
if err := session.HangupCall(context.Background(), "reject-call"); err != nil {
t.Fatal(err)
}
if !strings.HasPrefix(string(response), "SIP/2.0 486 Busy Here") {
t.Fatalf("reject response = %q", response)
}
calls := session.Calls()
if len(calls) != 1 || calls[0].State != "ended" || calls[0].EndedAt == nil {
t.Fatalf("terminal call status = %#v", calls)
}
}
func TestRejectedOutgoingCallRetainsSIPReason(t *testing.T) {
session := &Session{calls: make(map[string]*imsCall)}
call := &imsCall{public: vowifi.Call{ID: "rejected", State: "dialing"}}
session.calls[call.public.ID] = call
session.finishCall(call.public.ID, "failed", 484, "Address Incomplete\r\nignored")
calls := session.Calls()
if len(calls) != 1 || calls[0].State != "failed" || calls[0].SIPCode != 484 ||
calls[0].Reason != "Address Incomplete ignored" || calls[0].EndedAt == nil {
t.Fatalf("rejected call = %#v", calls)
}
}
func TestValidCallNumber(t *testing.T) {
if !validCallNumber("+447700900000") || validCallNumber("12\r\nBYE") {
t.Fatal("call number validation mismatch")
}
}

Some files were not shown because too many files have changed in this diff Show More