diff --git a/wavesrv/cmd/main-server.go b/wavesrv/cmd/main-server.go new file mode 100644 index 00000000..69ebfb1c --- /dev/null +++ b/wavesrv/cmd/main-server.go @@ -0,0 +1,887 @@ +package main + +import ( + "context" + "encoding/base64" + "encoding/json" + "errors" + "fmt" + "io" + "io/fs" + "log" + "mime/multipart" + "net/http" + "os" + "os/signal" + "path/filepath" + "regexp" + "runtime/debug" + "strconv" + "strings" + "sync" + "syscall" + "time" + + "github.com/google/uuid" + "github.com/gorilla/mux" + + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/server" + "github.com/commandlinedev/prompt-server/pkg/cmdrunner" + "github.com/commandlinedev/prompt-server/pkg/pcloud" + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/rtnstate" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/commandlinedev/prompt-server/pkg/scpacket" + "github.com/commandlinedev/prompt-server/pkg/scws" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/commandlinedev/prompt-server/pkg/wsshell" +) + +type WebFnType = func(http.ResponseWriter, *http.Request) + +const HttpReadTimeout = 5 * time.Second +const HttpWriteTimeout = 21 * time.Second +const HttpMaxHeaderBytes = 60000 +const HttpTimeoutDuration = 21 * time.Second + +const MainServerAddr = "127.0.0.1:1619" // PromptServer, P=16, S=19, PS=1619 +const WebSocketServerAddr = "127.0.0.1:1623" // PromptWebsock, P=16, W=23, PW=1623 +const MainServerDevAddr = "127.0.0.1:8090" +const WebSocketServerDevAddr = "127.0.0.1:8091" +const WSStateReconnectTime = 30 * time.Second +const WSStatePacketChSize = 20 + +const InitialTelemetryWait = 30 * time.Second +const TelemetryTick = 30 * time.Minute +const TelemetryInterval = 8 * time.Hour + +const MaxWriteFileMemSize = 20 * (1024 * 1024) // 20M + +var GlobalLock = &sync.Mutex{} +var WSStateMap = make(map[string]*scws.WSState) // clientid -> WsState +var GlobalAuthKey string +var BuildTime = "0" +var shutdownOnce sync.Once +var ContentTypeHeaderValidRe = regexp.MustCompile(`^\w+/[\w.+-]+$`) + +type ClientActiveState struct { + Fg bool `json:"fg"` + Active bool `json:"active"` + Open bool `json:"open"` +} + +func setWSState(state *scws.WSState) { + GlobalLock.Lock() + defer GlobalLock.Unlock() + WSStateMap[state.ClientId] = state +} + +func getWSState(clientId string) *scws.WSState { + GlobalLock.Lock() + defer GlobalLock.Unlock() + return WSStateMap[clientId] +} + +func removeWSStateAfterTimeout(clientId string, connectTime time.Time, waitDuration time.Duration) { + go func() { + time.Sleep(waitDuration) + GlobalLock.Lock() + defer GlobalLock.Unlock() + state := WSStateMap[clientId] + if state == nil || state.ConnectTime != connectTime { + return + } + delete(WSStateMap, clientId) + state.UnWatchScreen() + }() +} + +func HandleWs(w http.ResponseWriter, r *http.Request) { + shell, err := wsshell.StartWS(w, r) + if err != nil { + log.Printf("WebSocket Upgrade Failed %T: %v\n", w, err) + return + } + defer shell.Conn.Close() + clientId := r.URL.Query().Get("clientid") + if clientId == "" { + close(shell.WriteChan) + return + } + state := getWSState(clientId) + if state == nil { + state = scws.MakeWSState(clientId, GlobalAuthKey) + state.ReplaceShell(shell) + setWSState(state) + } else { + state.UpdateConnectTime() + state.ReplaceShell(shell) + } + stateConnectTime := state.GetConnectTime() + defer func() { + removeWSStateAfterTimeout(clientId, stateConnectTime, WSStateReconnectTime) + }() + log.Printf("WebSocket opened %s %s\n", state.ClientId, shell.RemoteAddr) + state.RunWSRead() +} + +// todo: sync multiple writes to the same fifoName into a single go-routine and do liveness checking on fifo +// if this returns an error, likely the fifo is dead and the cmd should be marked as 'done' +func writeToFifo(fifoName string, data []byte) error { + rwfd, err := os.OpenFile(fifoName, os.O_RDWR, 0600) + if err != nil { + return err + } + defer rwfd.Close() + fifoWriter, err := os.OpenFile(fifoName, os.O_WRONLY, 0600) // blocking open (open won't block because of rwfd) + if err != nil { + return err + } + defer fifoWriter.Close() + // this *could* block if the fifo buffer is full + // unlikely because if the reader is dead, and len(data) < pipe size, then the buffer will be empty and will clear after rwfd is closed + _, err = fifoWriter.Write(data) + if err != nil { + return err + } + return nil +} + +func HandleGetClientData(w http.ResponseWriter, r *http.Request) { + cdata, err := sstore.EnsureClientData(r.Context()) + if err != nil { + WriteJsonError(w, err) + return + } + cdata = cdata.Clean() + WriteJsonSuccess(w, cdata) + return +} + +func HandleSetWinSize(w http.ResponseWriter, r *http.Request) { + decoder := json.NewDecoder(r.Body) + var winSize sstore.ClientWinSizeType + err := decoder.Decode(&winSize) + if err != nil { + WriteJsonError(w, fmt.Errorf("error decoding json: %w", err)) + return + } + err = sstore.SetWinSize(r.Context(), winSize) + if err != nil { + WriteJsonError(w, fmt.Errorf("error setting winsize: %w", err)) + return + } + WriteJsonSuccess(w, true) + return +} + +// params: fg, active, open +func HandleLogActiveState(w http.ResponseWriter, r *http.Request) { + decoder := json.NewDecoder(r.Body) + var activeState ClientActiveState + err := decoder.Decode(&activeState) + if err != nil { + WriteJsonError(w, fmt.Errorf("error decoding json: %w", err)) + return + } + activity := sstore.ActivityUpdate{} + if activeState.Fg { + activity.FgMinutes = 1 + } + if activeState.Active { + activity.ActiveMinutes = 1 + } + if activeState.Open { + activity.OpenMinutes = 1 + } + activity.NumConns = remote.NumRemotes() + err = sstore.UpdateCurrentActivity(r.Context(), activity) + if err != nil { + WriteJsonError(w, fmt.Errorf("error updating activity: %w", err)) + return + } + WriteJsonSuccess(w, true) + return +} + +// params: screenid +func HandleGetScreenLines(w http.ResponseWriter, r *http.Request) { + qvals := r.URL.Query() + screenId := qvals.Get("screenid") + if _, err := uuid.Parse(screenId); err != nil { + WriteJsonError(w, fmt.Errorf("invalid screenid: %w", err)) + return + } + screenLines, err := sstore.GetScreenLinesById(r.Context(), screenId) + if err != nil { + WriteJsonError(w, err) + return + } + WriteJsonSuccess(w, screenLines) + return +} + +func HandleRtnState(w http.ResponseWriter, r *http.Request) { + defer func() { + r := recover() + if r == nil { + return + } + log.Printf("[error] in handlertnstate: %v\n", r) + debug.PrintStack() + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("panic: %v", r))) + return + }() + qvals := r.URL.Query() + screenId := qvals.Get("screenid") + lineId := qvals.Get("lineid") + if screenId == "" || lineId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify screenid and lineid"))) + return + } + if _, err := uuid.Parse(screenId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid screenid: %v", err))) + return + } + if _, err := uuid.Parse(lineId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid lineid: %v", err))) + return + } + data, err := rtnstate.GetRtnStateDiff(r.Context(), screenId, lineId) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("cannot get rtnstate diff: %v", err))) + return + } + w.WriteHeader(http.StatusOK) + w.Write(data) + return +} + +func HandleRemotePty(w http.ResponseWriter, r *http.Request) { + qvals := r.URL.Query() + remoteId := qvals.Get("remoteid") + if remoteId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify remoteid"))) + return + } + if _, err := uuid.Parse(remoteId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid remoteid: %v", err))) + return + } + realOffset, data, err := remote.ReadRemotePty(r.Context(), remoteId) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error reading ptyout file: %v", err))) + return + } + w.Header().Set("X-PtyDataOffset", strconv.FormatInt(realOffset, 10)) + w.WriteHeader(http.StatusOK) + w.Write(data) + return +} + +func HandleGetPtyOut(w http.ResponseWriter, r *http.Request) { + qvals := r.URL.Query() + screenId := qvals.Get("screenid") + lineId := qvals.Get("lineid") + if screenId == "" || lineId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify screenid and lineid"))) + return + } + if _, err := uuid.Parse(screenId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid screenid: %v", err))) + return + } + if _, err := uuid.Parse(lineId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid lineid: %v", err))) + return + } + realOffset, data, err := sstore.ReadFullPtyOutFile(r.Context(), screenId, lineId) + if err != nil { + if errors.Is(err, fs.ErrNotExist) { + w.WriteHeader(http.StatusOK) + return + } + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error reading ptyout file: %v", err))) + return + } + w.Header().Set("X-PtyDataOffset", strconv.FormatInt(realOffset, 10)) + w.WriteHeader(http.StatusOK) + w.Write(data) +} + +type writeFileParamsType struct { + ScreenId string `json:"screenid"` + LineId string `json:"lineid"` + Path string `json:"path"` + UseTemp bool `json:"usetemp,omitempty"` +} + +func parseWriteFileParams(r *http.Request) (*writeFileParamsType, multipart.File, error) { + err := r.ParseMultipartForm(MaxWriteFileMemSize) + if err != nil { + return nil, nil, fmt.Errorf("cannot parse multipart form data: %v", err) + } + form := r.MultipartForm + if len(form.Value["params"]) == 0 { + return nil, nil, fmt.Errorf("no params found") + } + paramsStr := form.Value["params"][0] + var params writeFileParamsType + err = json.Unmarshal([]byte(paramsStr), ¶ms) + if err != nil { + return nil, nil, fmt.Errorf("bad params json: %v", err) + } + if len(form.File["data"]) == 0 { + return nil, nil, fmt.Errorf("no data found") + } + fileHeader := form.File["data"][0] + file, err := fileHeader.Open() + if err != nil { + return nil, nil, fmt.Errorf("error opening multipart data file: %v", err) + } + return ¶ms, file, nil +} + +func HandleWriteFile(w http.ResponseWriter, r *http.Request) { + defer func() { + r := recover() + if r == nil { + return + } + log.Printf("[error] in write-file: %v\n", r) + debug.PrintStack() + WriteJsonError(w, fmt.Errorf("panic: %v", r)) + return + }() + w.Header().Set("Cache-Control", "no-cache") + params, mpFile, err := parseWriteFileParams(r) + if err != nil { + WriteJsonError(w, fmt.Errorf("error parsing multipart form params: %w", err)) + return + } + if params.ScreenId == "" || params.LineId == "" || params.Path == "" { + WriteJsonError(w, fmt.Errorf("invalid params, must set screenid, lineid, and path")) + return + } + if _, err := uuid.Parse(params.ScreenId); err != nil { + WriteJsonError(w, fmt.Errorf("invalid screenid: %v", err)) + return + } + if _, err := uuid.Parse(params.LineId); err != nil { + WriteJsonError(w, fmt.Errorf("invalid lineid: %v", err)) + return + } + _, cmd, err := sstore.GetLineCmdByLineId(r.Context(), params.ScreenId, params.LineId) + if err != nil { + WriteJsonError(w, fmt.Errorf("cannot retrieve line/cmd: %v", err)) + return + } + if cmd == nil { + WriteJsonError(w, fmt.Errorf("line not found")) + return + } + if cmd.Remote.RemoteId == "" { + WriteJsonError(w, fmt.Errorf("invalid line, no remote")) + return + } + msh := remote.GetRemoteById(cmd.Remote.RemoteId) + if msh == nil { + WriteJsonError(w, fmt.Errorf("invalid line, cannot resolve remote")) + return + } + rrState := msh.GetRemoteRuntimeState() + fullPath, err := rrState.ExpandHomeDir(params.Path) + if err != nil { + WriteJsonError(w, fmt.Errorf("error expanding homedir: %v", err)) + return + } + cwd := cmd.FeState["cwd"] + writePk := packet.MakeWriteFilePacket() + writePk.ReqId = uuid.New().String() + writePk.UseTemp = params.UseTemp + if filepath.IsAbs(fullPath) { + writePk.Path = fullPath + } else { + writePk.Path = filepath.Join(cwd, fullPath) + } + iter, err := msh.PacketRpcIter(r.Context(), writePk) + if err != nil { + WriteJsonError(w, fmt.Errorf("error: %v", err)) + return + } + // first packet should be WriteFileReady + readyIf, err := iter.Next(r.Context()) + if err != nil { + WriteJsonError(w, fmt.Errorf("error while getting ready response: %w", err)) + return + } + readyPk, ok := readyIf.(*packet.WriteFileReadyPacketType) + if !ok { + WriteJsonError(w, fmt.Errorf("bad ready packet received: %T", readyIf)) + return + } + if readyPk.Error != "" { + WriteJsonError(w, fmt.Errorf("ready error: %s", readyPk.Error)) + return + } + var buffer [server.MaxFileDataPacketSize]byte + bufSlice := buffer[:] + for { + dataPk := packet.MakeFileDataPacket(writePk.ReqId) + nr, err := io.ReadFull(mpFile, bufSlice) + if err == io.ErrUnexpectedEOF || err == io.EOF { + dataPk.Eof = true + } else if err != nil { + dataErr := fmt.Errorf("error reading file data: %v", err) + dataPk.Error = dataErr.Error() + msh.SendFileData(dataPk) + WriteJsonError(w, dataErr) + return + } + if nr > 0 { + dataPk.Data = make([]byte, nr) + copy(dataPk.Data, bufSlice[0:nr]) + } + msh.SendFileData(dataPk) + if dataPk.Eof { + break + } + // slight throttle for sending packets + time.Sleep(10 * time.Millisecond) + } + doneIf, err := iter.Next(r.Context()) + if err != nil { + WriteJsonError(w, fmt.Errorf("error while getting done response: %w", err)) + return + } + donePk, ok := doneIf.(*packet.WriteFileDonePacketType) + if !ok { + WriteJsonError(w, fmt.Errorf("bad done packet received: %T", doneIf)) + return + } + if donePk.Error != "" { + WriteJsonError(w, fmt.Errorf("dne error: %s", donePk.Error)) + return + } + WriteJsonSuccess(w, nil) + return +} + +func HandleReadFile(w http.ResponseWriter, r *http.Request) { + qvals := r.URL.Query() + screenId := qvals.Get("screenid") + lineId := qvals.Get("lineid") + path := qvals.Get("path") // validate path? + contentType := qvals.Get("mimetype") + if contentType == "" { + contentType = "application/octet-stream" + } + if screenId == "" || lineId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify sessionid, screenid, and lineid"))) + return + } + if path == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("must specify path"))) + return + } + if _, err := uuid.Parse(screenId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid screenid: %v", err))) + return + } + if _, err := uuid.Parse(lineId); err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid lineid: %v", err))) + return + } + if !ContentTypeHeaderValidRe.MatchString(contentType) { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid mimetype specified"))) + return + } + _, cmd, err := sstore.GetLineCmdByLineId(r.Context(), screenId, lineId) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid lineid: %v", err))) + return + } + if cmd == nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid line, no cmd"))) + return + } + if cmd.Remote.RemoteId == "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid line, no remote"))) + return + } + msh := remote.GetRemoteById(cmd.Remote.RemoteId) + if msh == nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("invalid line, cannot resolve remote"))) + return + } + rrState := msh.GetRemoteRuntimeState() + fullPath, err := rrState.ExpandHomeDir(path) + if err != nil { + WriteJsonError(w, fmt.Errorf("error expanding homedir: %v", err)) + return + } + streamPk := packet.MakeStreamFilePacket() + streamPk.ReqId = uuid.New().String() + cwd := cmd.FeState["cwd"] + if filepath.IsAbs(fullPath) { + streamPk.Path = fullPath + } else { + streamPk.Path = filepath.Join(cwd, fullPath) + } + iter, err := msh.StreamFile(r.Context(), streamPk) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error trying to stream file: %v", err))) + return + } + defer iter.Close() + respIf, err := iter.Next(r.Context()) + if err != nil { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error getting streamfile response: %v", err))) + return + } + resp, ok := respIf.(*packet.StreamFileResponseType) + if !ok { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("bad response packet type: %T", respIf))) + return + } + if resp.Error != "" { + w.WriteHeader(500) + w.Write([]byte(fmt.Sprintf("error response: %s", resp.Error))) + return + } + infoJson, _ := json.Marshal(resp.Info) + w.Header().Set("X-FileInfo", base64.StdEncoding.EncodeToString(infoJson)) + w.Header().Set("Content-Type", contentType) + w.WriteHeader(http.StatusOK) + for { + dataPkIf, err := iter.Next(r.Context()) + if err != nil { + log.Printf("error in read-file while getting data: %v\n", err) + break + } + if dataPkIf == nil { + break + } + dataPk, ok := dataPkIf.(*packet.FileDataPacketType) + if !ok { + log.Printf("error in read-file, invalid data packet type: %T", dataPkIf) + break + } + if dataPk.Error != "" { + log.Printf("in read-file, data packet error: %s", dataPk.Error) + break + } + w.Write(dataPk.Data) + } + return +} + +func WriteJsonError(w http.ResponseWriter, errVal error) { + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(200) + errMap := make(map[string]interface{}) + errMap["error"] = errVal.Error() + barr, _ := json.Marshal(errMap) + w.Write(barr) + return +} + +func WriteJsonSuccess(w http.ResponseWriter, data interface{}) { + w.Header().Set("Content-Type", "application/json") + rtnMap := make(map[string]interface{}) + rtnMap["success"] = true + if data != nil { + rtnMap["data"] = data + } + barr, err := json.Marshal(rtnMap) + if err != nil { + WriteJsonError(w, err) + return + } + w.WriteHeader(200) + w.Write(barr) + return +} + +func HandleRunCommand(w http.ResponseWriter, r *http.Request) { + defer func() { + r := recover() + if r == nil { + return + } + log.Printf("[error] in run-command: %v\n", r) + debug.PrintStack() + WriteJsonError(w, fmt.Errorf("panic: %v", r)) + return + }() + w.Header().Set("Cache-Control", "no-cache") + decoder := json.NewDecoder(r.Body) + var commandPk scpacket.FeCommandPacketType + err := decoder.Decode(&commandPk) + if err != nil { + WriteJsonError(w, fmt.Errorf("error decoding json: %w", err)) + return + } + update, err := cmdrunner.HandleCommand(r.Context(), &commandPk) + if err != nil { + WriteJsonError(w, err) + return + } + if update != nil { + update.Clean() + } + WriteJsonSuccess(w, update) + return +} + +func AuthKeyWrap(fn WebFnType) WebFnType { + return func(w http.ResponseWriter, r *http.Request) { + reqAuthKey := r.Header.Get("X-AuthKey") + if reqAuthKey == "" { + w.WriteHeader(500) + w.Write([]byte("no x-authkey header")) + return + } + if reqAuthKey != GlobalAuthKey { + w.WriteHeader(500) + w.Write([]byte("x-authkey header is invalid")) + return + } + w.Header().Set("Cache-Control", "no-cache") + fn(w, r) + } +} + +func runWebSocketServer() { + gr := mux.NewRouter() + gr.HandleFunc("/ws", HandleWs) + serverAddr := WebSocketServerAddr + if scbase.IsDevMode() { + serverAddr = WebSocketServerDevAddr + } + server := &http.Server{ + Addr: serverAddr, + ReadTimeout: HttpReadTimeout, + WriteTimeout: HttpWriteTimeout, + MaxHeaderBytes: HttpMaxHeaderBytes, + Handler: gr, + } + server.SetKeepAlivesEnabled(false) + log.Printf("Running websocket server on %s\n", serverAddr) + err := server.ListenAndServe() + if err != nil { + log.Printf("[error] trying to run websocket server: %v\n", err) + } +} + +func test() error { + return nil +} + +func sendTelemetryWrapper() { + defer func() { + r := recover() + if r == nil { + return + } + log.Printf("[error] in sendTelemetryWrapper: %v\n", r) + debug.PrintStack() + return + }() + ctx, cancelFn := context.WithTimeout(context.Background(), 5*time.Second) + defer cancelFn() + err := pcloud.SendTelemetry(ctx, false) + if err != nil { + log.Printf("[error] sending telemetry: %v\n", err) + } +} + +func telemetryLoop() { + var lastSent time.Time + time.Sleep(InitialTelemetryWait) + for { + dur := time.Now().Sub(lastSent) + if lastSent.IsZero() || dur >= TelemetryInterval { + lastSent = time.Now() + sendTelemetryWrapper() + } + time.Sleep(TelemetryTick) + } +} + +// watch stdin, kill server if stdin is closed +func stdinReadWatch() { + buf := make([]byte, 1024) + for { + _, err := os.Stdin.Read(buf) + if err != nil { + doShutdown(fmt.Sprintf("stdin closed/error (%v)", err)) + break + } + } +} + +// ignore SIGHUP +func installSignalHandlers() { + sigCh := make(chan os.Signal, 1) + signal.Notify(sigCh, syscall.SIGHUP) + go func() { + for sig := range sigCh { + doShutdown(fmt.Sprintf("got signal %v", sig)) + break + } + }() +} + +func doShutdown(reason string) { + shutdownOnce.Do(func() { + log.Printf("[prompt] local server %v, start shutdown\n", reason) + sendTelemetryWrapper() + log.Printf("[prompt] closing db connection\n") + sstore.CloseDB() + log.Printf("[prompt] *** shutting down local server\n") + time.Sleep(1 * time.Second) + syscall.Kill(syscall.Getpid(), syscall.SIGINT) + time.Sleep(5 * time.Second) + syscall.Kill(syscall.Getpid(), syscall.SIGKILL) + }) +} + +func main() { + scbase.BuildTime = BuildTime + + if len(os.Args) >= 2 && os.Args[1] == "--test" { + log.Printf("running test fn\n") + err := test() + if err != nil { + log.Printf("[error] %v\n", err) + } + return + } + + scHomeDir := scbase.GetPromptHomeDir() + log.Printf("[prompt] *** starting local server\n") + log.Printf("[prompt] local server version %s+%s\n", scbase.PromptVersion, scbase.BuildTime) + log.Printf("[prompt] homedir = %q\n", scHomeDir) + + scLock, err := scbase.AcquirePromptLock() + if err != nil || scLock == nil { + log.Printf("[error] cannot acquire prompt lock: %v\n", err) + return + } + if len(os.Args) >= 2 && strings.HasPrefix(os.Args[1], "--migrate") { + err := sstore.MigrateCommandOpts(os.Args[1:]) + if err != nil { + log.Printf("[error] migrate cmd: %v\n", err) + } + return + } + authKey, err := scbase.ReadPromptAuthKey() + if err != nil { + log.Printf("[error] %v\n", err) + return + } + GlobalAuthKey = authKey + err = sstore.TryMigrateUp() + if err != nil { + log.Printf("[error] migrate up: %v\n", err) + return + } + clientData, err := sstore.EnsureClientData(context.Background()) + if err != nil { + log.Printf("[error] ensuring client data: %v\n", err) + return + } + log.Printf("userid = %s\n", clientData.UserId) + err = sstore.EnsureLocalRemote(context.Background()) + if err != nil { + log.Printf("[error] ensuring local remote: %v\n", err) + return + } + _, err = sstore.EnsureDefaultSession(context.Background()) + if err != nil { + log.Printf("[error] ensuring default session: %v\n", err) + return + } + err = remote.LoadRemotes(context.Background()) + if err != nil { + log.Printf("[error] loading remotes: %v\n", err) + return + } + + err = sstore.HangupAllRunningCmds(context.Background()) + if err != nil { + log.Printf("[error] calling HUP on all running commands: %v\n", err) + } + err = sstore.ReInitFocus(context.Background()) + if err != nil { + log.Printf("[error] resetting screen focus: %v\n", err) + } + + log.Printf("PCLOUD_ENDPOINT=%s\n", pcloud.GetEndpoint()) + err = sstore.UpdateCurrentActivity(context.Background(), sstore.ActivityUpdate{NumConns: remote.NumRemotes()}) // set at least one record into activity + if err != nil { + log.Printf("[error] updating activity: %v\n", err) + } + installSignalHandlers() + go telemetryLoop() + go stdinReadWatch() + go runWebSocketServer() + go func() { + time.Sleep(10 * time.Second) + pcloud.StartUpdateWriter() + }() + gr := mux.NewRouter() + gr.HandleFunc("/api/ptyout", AuthKeyWrap(HandleGetPtyOut)) + gr.HandleFunc("/api/remote-pty", AuthKeyWrap(HandleRemotePty)) + gr.HandleFunc("/api/rtnstate", AuthKeyWrap(HandleRtnState)) + gr.HandleFunc("/api/get-screen-lines", AuthKeyWrap(HandleGetScreenLines)) + gr.HandleFunc("/api/run-command", AuthKeyWrap(HandleRunCommand)).Methods("POST") + gr.HandleFunc("/api/get-client-data", AuthKeyWrap(HandleGetClientData)) + gr.HandleFunc("/api/set-winsize", AuthKeyWrap(HandleSetWinSize)) + gr.HandleFunc("/api/log-active-state", AuthKeyWrap(HandleLogActiveState)) + gr.HandleFunc("/api/read-file", AuthKeyWrap(HandleReadFile)) + gr.HandleFunc("/api/write-file", AuthKeyWrap(HandleWriteFile)).Methods("POST") + serverAddr := MainServerAddr + if scbase.IsDevMode() { + serverAddr = MainServerDevAddr + } + server := &http.Server{ + Addr: serverAddr, + ReadTimeout: HttpReadTimeout, + WriteTimeout: HttpWriteTimeout, + MaxHeaderBytes: HttpMaxHeaderBytes, + Handler: http.TimeoutHandler(gr, HttpTimeoutDuration, "Timeout"), + } + server.SetKeepAlivesEnabled(false) + log.Printf("Running main server on %s\n", serverAddr) + err = server.ListenAndServe() + if err != nil { + log.Printf("ERROR: %v\n", err) + } +} diff --git a/wavesrv/db/db.go b/wavesrv/db/db.go new file mode 100644 index 00000000..de43f7c6 --- /dev/null +++ b/wavesrv/db/db.go @@ -0,0 +1,9 @@ +// provides the io/fs for DB migrations +package db + +import "embed" + +// since embeds must be relative to the package directory, this source file is required + +//go:embed migrations/*.sql +var MigrationFS embed.FS diff --git a/wavesrv/db/migrations/000001_init.down.sql b/wavesrv/db/migrations/000001_init.down.sql new file mode 100644 index 00000000..4892a38e --- /dev/null +++ b/wavesrv/db/migrations/000001_init.down.sql @@ -0,0 +1,13 @@ +DROP TABLE client; +DROP TABLE session; +DROP TABLE window; +DROP TABLE screen; +DROP TABLE screen_window; +DROP TABLE remote_instance; +DROP TABLE line; +DROP TABLE remote; +DROP TABLE cmd; +DROP TABLE history; +DROP TABLE state_base; +DROP TABLE state_diff; + diff --git a/wavesrv/db/migrations/000001_init.up.sql b/wavesrv/db/migrations/000001_init.up.sql new file mode 100644 index 00000000..e13caed8 --- /dev/null +++ b/wavesrv/db/migrations/000001_init.up.sql @@ -0,0 +1,167 @@ +CREATE TABLE client ( + clientid varchar(36) NOT NULL, + userid varchar(36) NOT NULL, + activesessionid varchar(36) NOT NULL, + userpublickeybytes blob NOT NULL, + userprivatekeybytes blob NOT NULL, + winsize json NOT NULL +); + +CREATE TABLE session ( + sessionid varchar(36) PRIMARY KEY, + name varchar(50) NOT NULL, + sessionidx int NOT NULL, + activescreenid varchar(36) NOT NULL, + notifynum int NOT NULL, + archived boolean NOT NULL, + archivedts bigint NOT NULL, + ownerid varchar(36) NOT NULL, + sharemode varchar(12) NOT NULL, + accesskey varchar(36) NOT NULL +); + +CREATE TABLE window ( + sessionid varchar(36) NOT NULL, + windowid varchar(36) NOT NULL, + curremoteownerid varchar(36) NOT NULL, + curremoteid varchar(36) NOT NULL, + curremotename varchar(50) NOT NULL, + nextlinenum int NOT NULL, + winopts json NOT NULL, + ownerid varchar(36) NOT NULL, + sharemode varchar(12) NOT NULL, + shareopts json NOT NULL, + PRIMARY KEY (sessionid, windowid) +); + +CREATE TABLE screen ( + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + name varchar(50) NOT NULL, + activewindowid varchar(36) NOT NULL, + screenidx int NOT NULL, + screenopts json NOT NULL, + ownerid varchar(36) NOT NULL, + sharemode varchar(12) NOT NULL, + incognito boolean NOT NULL, + archived boolean NOT NULL, + archivedts bigint NOT NULL, + PRIMARY KEY (sessionid, screenid) +); + +CREATE TABLE screen_window ( + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + windowid varchar(36) NOT NULL, + name varchar(50) NOT NULL, + layout json NOT NULL, + selectedline int NOT NULL, + anchor json NOT NULL, + focustype varchar(12) NOT NULL, + PRIMARY KEY (sessionid, screenid, windowid) +); + +CREATE TABLE remote_instance ( + riid varchar(36) PRIMARY KEY, + name varchar(50) NOT NULL, + sessionid varchar(36) NOT NULL, + windowid varchar(36) NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + festate json NOT NULL, + statebasehash varchar(36) NOT NULL, + statediffhasharr json NOT NULL +); + +CREATE TABLE state_base ( + basehash varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + version varchar(200) NOT NULL, + data blob NOT NULL +); + +CREATE TABLE state_diff ( + diffhash varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + basehash varchar(36) NOT NULL, + diffhasharr json NOT NULL, + data blob NOT NULL +); + +CREATE TABLE line ( + sessionid varchar(36) NOT NULL, + windowid varchar(36) NOT NULL, + userid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + ts bigint NOT NULL, + linenum int NOT NULL, + linenumtemp boolean NOT NULL, + linetype varchar(10) NOT NULL, + linelocal boolean NOT NULL, + text text NOT NULL, + cmdid varchar(36) NOT NULL, + ephemeral boolean NOT NULL, + contentheight int NOT NULL, + star int NOT NULL, + archived boolean NOT NULL, + PRIMARY KEY (sessionid, windowid, lineid) +); + +CREATE TABLE remote ( + remoteid varchar(36) PRIMARY KEY, + physicalid varchar(36) NOT NULL, + remotetype varchar(10) NOT NULL, + remotealias varchar(50) NOT NULL, + remotecanonicalname varchar(200) NOT NULL, + remotesudo boolean NOT NULL, + remoteuser varchar(50) NOT NULL, + remotehost varchar(200) NOT NULL, + connectmode varchar(20) NOT NULL, + autoinstall boolean NOT NULL, + sshopts json NOT NULL, + remoteopts json NOT NULL, + lastconnectts bigint NOT NULL, + local boolean NOT NULL, + archived boolean NOT NULL, + remoteidx int NOT NULL +); + +CREATE TABLE cmd ( + sessionid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + remotename varchar(50) NOT NULL, + cmdstr text NOT NULL, + festate json NOT NULL, + statebasehash varchar(36) NOT NULL, + statediffhasharr json NOT NULL, + termopts json NOT NULL, + origtermopts json NOT NULL, + status varchar(10) NOT NULL, + startpk json NOT NULL, + doneinfo json NOT NULL, + runout json NOT NULL, + rtnstate boolean NOT NULL, + rtnbasehash varchar(36) NOT NULL, + rtndiffhasharr json NOT NULL, + PRIMARY KEY (sessionid, cmdid) +); + +CREATE TABLE history ( + historyid varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + userid varchar(36) NOT NULL, + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + windowid varchar(36) NOT NULL, + lineid int NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + remotename varchar(50) NOT NULL, + haderror boolean NOT NULL, + cmdid varchar(36) NOT NULL, + cmdstr text NOT NULL, + ismetacmd boolean, + incognito boolean +); diff --git a/wavesrv/db/migrations/000002_activity.down.sql b/wavesrv/db/migrations/000002_activity.down.sql new file mode 100644 index 00000000..75ccae38 --- /dev/null +++ b/wavesrv/db/migrations/000002_activity.down.sql @@ -0,0 +1,3 @@ +DROP TABLE activity; + +ALTER TABLE client DROP COLUMN clientopts; diff --git a/wavesrv/db/migrations/000002_activity.up.sql b/wavesrv/db/migrations/000002_activity.up.sql new file mode 100644 index 00000000..d6a84aa3 --- /dev/null +++ b/wavesrv/db/migrations/000002_activity.up.sql @@ -0,0 +1,11 @@ +CREATE TABLE activity ( + day varchar(20) PRIMARY KEY, + uploaded boolean NOT NULL, + tdata json NOT NULL, + tzname varchar(50) NOT NULL, + tzoffset int NOT NULL, + clientversion varchar(20) NOT NULL, + clientarch varchar(20) NOT NULL +); + +ALTER TABLE client ADD COLUMN clientopts json NOT NULL DEFAULT ''; diff --git a/wavesrv/db/migrations/000003_renderer.down.sql b/wavesrv/db/migrations/000003_renderer.down.sql new file mode 100644 index 00000000..4ef6ef5a --- /dev/null +++ b/wavesrv/db/migrations/000003_renderer.down.sql @@ -0,0 +1,2 @@ +ALTER TABLE line DROP COLUMN renderer; + diff --git a/wavesrv/db/migrations/000003_renderer.up.sql b/wavesrv/db/migrations/000003_renderer.up.sql new file mode 100644 index 00000000..43f9de43 --- /dev/null +++ b/wavesrv/db/migrations/000003_renderer.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE line ADD COLUMN renderer varchar(50) NOT NULL DEFAULT ''; + diff --git a/wavesrv/db/migrations/000004_bookmarks.down.sql b/wavesrv/db/migrations/000004_bookmarks.down.sql new file mode 100644 index 00000000..13248684 --- /dev/null +++ b/wavesrv/db/migrations/000004_bookmarks.down.sql @@ -0,0 +1,6 @@ +DROP TABLE bookmark; +DROP TABLE bookmark_order; +DROP TABLE bookmark_cmd; + +ALTER TABLE line DROP COLUMN bookmarked; +ALTER TABLE line DROP COLUMN pinned; diff --git a/wavesrv/db/migrations/000004_bookmarks.up.sql b/wavesrv/db/migrations/000004_bookmarks.up.sql new file mode 100644 index 00000000..35d3589e --- /dev/null +++ b/wavesrv/db/migrations/000004_bookmarks.up.sql @@ -0,0 +1,26 @@ +CREATE TABLE bookmark ( + bookmarkid varchar(36) PRIMARY KEY, + createdts bigint NOT NULL, + cmdstr text NOT NULL, + alias varchar(50) NOT NULL, + tags json NOT NULL, + description text NOT NULL +); + +CREATE TABLE bookmark_order ( + tag varchar(50) NOT NULL, + bookmarkid varchar(36) NOT NULL, + orderidx int NOT NULL, + PRIMARY KEY (tag, bookmarkid) +); + +CREATE TABLE bookmark_cmd ( + bookmarkid varchar(36) NOT NULL, + sessionid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + PRIMARY KEY (bookmarkid, sessionid, cmdid) +); + +ALTER TABLE line ADD COLUMN bookmarked boolean NOT NULL DEFAULT 0; +ALTER TABLE line ADD COLUMN pinned boolean NOT NULL DEFAULT 0; + diff --git a/wavesrv/db/migrations/000005_buildtime.down.sql b/wavesrv/db/migrations/000005_buildtime.down.sql new file mode 100644 index 00000000..2c83f8ef --- /dev/null +++ b/wavesrv/db/migrations/000005_buildtime.down.sql @@ -0,0 +1,2 @@ +ALTER TABLE activity DROP COLUMN buildtime; +ALTER TABLE activity DROP COLUMN osrelease; diff --git a/wavesrv/db/migrations/000005_buildtime.up.sql b/wavesrv/db/migrations/000005_buildtime.up.sql new file mode 100644 index 00000000..1006efa1 --- /dev/null +++ b/wavesrv/db/migrations/000005_buildtime.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE activity ADD COLUMN buildtime varchar(20) NOT NULL DEFAULT '-'; +ALTER TABLE activity ADD COLUMN osrelease varchar(20) NOT NULL DEFAULT '-'; diff --git a/wavesrv/db/migrations/000006_feopts.down.sql b/wavesrv/db/migrations/000006_feopts.down.sql new file mode 100644 index 00000000..7f9b0092 --- /dev/null +++ b/wavesrv/db/migrations/000006_feopts.down.sql @@ -0,0 +1 @@ +ALTER TABLE client DROP COLUMN feopts; diff --git a/wavesrv/db/migrations/000006_feopts.up.sql b/wavesrv/db/migrations/000006_feopts.up.sql new file mode 100644 index 00000000..f822f566 --- /dev/null +++ b/wavesrv/db/migrations/000006_feopts.up.sql @@ -0,0 +1,3 @@ +ALTER TABLE client ADD COLUMN feopts json NOT NULL DEFAULT '{}'; + + diff --git a/wavesrv/db/migrations/000007_playbooks.down.sql b/wavesrv/db/migrations/000007_playbooks.down.sql new file mode 100644 index 00000000..4ed8c116 --- /dev/null +++ b/wavesrv/db/migrations/000007_playbooks.down.sql @@ -0,0 +1,3 @@ +DROP TABLE playbook; + +DROP TABLE playbook_entry; diff --git a/wavesrv/db/migrations/000007_playbooks.up.sql b/wavesrv/db/migrations/000007_playbooks.up.sql new file mode 100644 index 00000000..506e4ea5 --- /dev/null +++ b/wavesrv/db/migrations/000007_playbooks.up.sql @@ -0,0 +1,16 @@ +CREATE TABLE playbook ( + playbookid varchar(36) PRIMARY KEY, + playbookname varchar(100) NOT NULL, + description text NOT NULL, + entryids json NOT NULL +); + +CREATE TABLE playbook_entry ( + entryid varchar(36) PRIMARY KEY, + playbookid varchar(36) NOT NULL, + description text NOT NULL, + alias varchar(50) NOT NULL, + cmdstr text NOT NULL, + createdts bigint NOT NULL, + updatedts bigint NOT NULL +); diff --git a/wavesrv/db/migrations/000008_cloudsession.down.sql b/wavesrv/db/migrations/000008_cloudsession.down.sql new file mode 100644 index 00000000..bf15963a --- /dev/null +++ b/wavesrv/db/migrations/000008_cloudsession.down.sql @@ -0,0 +1,5 @@ +ALTER TABLE session ADD COLUMN accesskey DEFAULT ''; +ALTER TABLE session ADD COLUMN ownerid DEFAULT ''; + +DROP TABLE cloud_session; +DROP TABLE cloud_update; diff --git a/wavesrv/db/migrations/000008_cloudsession.up.sql b/wavesrv/db/migrations/000008_cloudsession.up.sql new file mode 100644 index 00000000..860e61b0 --- /dev/null +++ b/wavesrv/db/migrations/000008_cloudsession.up.sql @@ -0,0 +1,20 @@ +ALTER TABLE session DROP COLUMN accesskey; +ALTER TABLE session DROP COLUMN ownerid; + +CREATE TABLE cloud_session ( + sessionid varchar(36) PRIMARY KEY, + viewkey varchar(50) NOT NULL, + writekey varchar(50) NOT NULL, + enckey varchar(100) NOT NULL, + enctype varchar(50) NOT NULL, + vts bigint NOT NULL, + acl json NOT NULL +); + +CREATE TABLE cloud_update ( + updateid varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + updatetype varchar(50) NOT NULL, + updatekeys json NOT NULL +); + diff --git a/wavesrv/db/migrations/000009_screenprimary.down.sql b/wavesrv/db/migrations/000009_screenprimary.down.sql new file mode 100644 index 00000000..a4b58677 --- /dev/null +++ b/wavesrv/db/migrations/000009_screenprimary.down.sql @@ -0,0 +1,3 @@ +-- invalid, will throw an error, cannot migrate down +SELECT x; + diff --git a/wavesrv/db/migrations/000009_screenprimary.up.sql b/wavesrv/db/migrations/000009_screenprimary.up.sql new file mode 100644 index 00000000..48efa4f8 --- /dev/null +++ b/wavesrv/db/migrations/000009_screenprimary.up.sql @@ -0,0 +1,56 @@ +CREATE TABLE new_screen ( + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + windowid varchar(36) NOT NULL, + name varchar(50) NOT NULL, + screenidx int NOT NULL, + screenopts json NOT NULL, + ownerid varchar(36) NOT NULL, + sharemode varchar(12) NOT NULL, + curremoteownerid varchar(36) NOT NULL, + curremoteid varchar(36) NOT NULL, + curremotename varchar(50) NOT NULL, + nextlinenum int NOT NULL, + selectedline int NOT NULL, + anchor json NOT NULL, + focustype varchar(12) NOT NULL, + archived boolean NOT NULL, + archivedts bigint NOT NULL, + PRIMARY KEY (sessionid, screenid) +); + +INSERT INTO new_screen +SELECT + s.sessionid, + s.screenid, + w.windowid, + s.name, + s.screenidx, + json_patch(s.screenopts, w.winopts), + s.ownerid, + s.sharemode, + w.curremoteownerid, + w.curremoteid, + w.curremotename, + w.nextlinenum, + sw.selectedline, + sw.anchor, + sw.focustype, + s.archived, + s.archivedts +FROM + screen s, + screen_window sw, + window w +WHERE + s.screenid = sw.screenid + AND sw.windowid = w.windowid +; + +DROP TABLE screen; +DROP TABLE screen_window; +DROP TABLE window; + +ALTER TABLE new_screen RENAME TO screen; + + diff --git a/wavesrv/db/migrations/000010_removewindowid.down.sql b/wavesrv/db/migrations/000010_removewindowid.down.sql new file mode 100644 index 00000000..6332dc5b --- /dev/null +++ b/wavesrv/db/migrations/000010_removewindowid.down.sql @@ -0,0 +1,2 @@ +-- invalid, will throw an error, cannot migrate down +SELECT x; diff --git a/wavesrv/db/migrations/000010_removewindowid.up.sql b/wavesrv/db/migrations/000010_removewindowid.up.sql new file mode 100644 index 00000000..10b1bbc0 --- /dev/null +++ b/wavesrv/db/migrations/000010_removewindowid.up.sql @@ -0,0 +1,17 @@ +ALTER TABLE remote_instance RENAME COLUMN windowid TO screenid; +ALTER TABLE line RENAME COLUMN windowid TO screenid; + +UPDATE remote_instance +SET screenid = COALESCE((SELECT screen.screenid FROM screen WHERE screen.windowid = remote_instance.screenid), '') +WHERE screenid <> '' +; + +UPDATE line +SET screenid = COALESCE((SELECT screen.screenid FROM screen WHERE screen.windowid = line.screenid), '') +WHERE screenid <> '' +; + +ALTER TABLE history DROP COLUMN windowid; +ALTER TABLE screen DROP COLUMN windowid; + + diff --git a/wavesrv/db/migrations/000011_cmdscreenid.down.sql b/wavesrv/db/migrations/000011_cmdscreenid.down.sql new file mode 100644 index 00000000..62fbbdf0 --- /dev/null +++ b/wavesrv/db/migrations/000011_cmdscreenid.down.sql @@ -0,0 +1,2 @@ +ALTER TABLE cmd DROP COLUMN screenid; + diff --git a/wavesrv/db/migrations/000011_cmdscreenid.up.sql b/wavesrv/db/migrations/000011_cmdscreenid.up.sql new file mode 100644 index 00000000..13ea21bd --- /dev/null +++ b/wavesrv/db/migrations/000011_cmdscreenid.up.sql @@ -0,0 +1,5 @@ +ALTER TABLE cmd ADD COLUMN screenid varchar(36) NOT NULL DEFAULT ''; + +UPDATE cmd +SET screenid = (SELECT line.screenid FROM line WHERE line.cmdid = cmd.cmdid) +; diff --git a/wavesrv/db/migrations/000012_historylinenum.down.sql b/wavesrv/db/migrations/000012_historylinenum.down.sql new file mode 100644 index 00000000..89bde3e6 --- /dev/null +++ b/wavesrv/db/migrations/000012_historylinenum.down.sql @@ -0,0 +1 @@ +ALTER TABLE history DROP COLUMN linenum; diff --git a/wavesrv/db/migrations/000012_historylinenum.up.sql b/wavesrv/db/migrations/000012_historylinenum.up.sql new file mode 100644 index 00000000..d7dc6353 --- /dev/null +++ b/wavesrv/db/migrations/000012_historylinenum.up.sql @@ -0,0 +1,6 @@ +ALTER TABLE history ADD COLUMN linenum int NOT NULL DEFAULT 0; + +UPDATE history +SET linenum = COALESCE((SELECT line.linenum FROM line WHERE line.lineid = history.lineid), 0) +; + diff --git a/wavesrv/db/migrations/000013_cmdmigration.down.sql b/wavesrv/db/migrations/000013_cmdmigration.down.sql new file mode 100644 index 00000000..6332dc5b --- /dev/null +++ b/wavesrv/db/migrations/000013_cmdmigration.down.sql @@ -0,0 +1,2 @@ +-- invalid, will throw an error, cannot migrate down +SELECT x; diff --git a/wavesrv/db/migrations/000013_cmdmigration.up.sql b/wavesrv/db/migrations/000013_cmdmigration.up.sql new file mode 100644 index 00000000..eed358ee --- /dev/null +++ b/wavesrv/db/migrations/000013_cmdmigration.up.sql @@ -0,0 +1,123 @@ +DELETE FROM cmd +WHERE screenid = ''; + +DELETE FROM line +WHERE screenid = ''; + +DELETE FROM cmd +WHERE cmdid NOT IN (SELECT cmdid FROM line); + +DELETE FROM line +WHERE cmdid <> '' AND cmdid NOT IN (SELECT cmdid FROM cmd); + +CREATE TABLE new_bookmark_cmd ( + bookmarkid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + PRIMARY KEY (bookmarkid, screenid, cmdid) +); +INSERT INTO new_bookmark_cmd +SELECT + b.bookmarkid, + c.screenid, + c.cmdid +FROM bookmark_cmd b, cmd c +WHERE b.cmdid = c.cmdid; +DROP TABLE bookmark_cmd; +ALTER TABLE new_bookmark_cmd RENAME TO bookmark_cmd; + +ALTER TABLE client ADD COLUMN cmdstoretype varchar(20) DEFAULT 'session'; + +CREATE TABLE cmd_migrate ( + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL +); +INSERT INTO cmd_migrate +SELECT sessionid, screenid, cmdid +FROM cmd; + +-- update primary key for screen +CREATE TABLE new_screen ( + screenid varchar(36) NOT NULL, + sessionid varchar(36) NOT NULL, + name varchar(50) NOT NULL, + screenidx int NOT NULL, + screenopts json NOT NULL, + ownerid varchar(36) NOT NULL, + sharemode varchar(12) NOT NULL, + curremoteownerid varchar(36) NOT NULL, + curremoteid varchar(36) NOT NULL, + curremotename varchar(50) NOT NULL, + nextlinenum int NOT NULL, + selectedline int NOT NULL, + anchor json NOT NULL, + focustype varchar(12) NOT NULL, + archived boolean NOT NULL, + archivedts bigint NOT NULL, + PRIMARY KEY (screenid) +); +INSERT INTO new_screen +SELECT screenid, sessionid, name, screenidx, screenopts, ownerid, sharemode, + curremoteownerid, curremoteid, curremotename, nextlinenum, selectedline, + anchor, focustype, archived, archivedts +FROM screen; +DROP TABLE screen; +ALTER TABLE new_screen RENAME TO screen; + +-- drop sessionid from line +CREATE TABLE new_line ( + screenid varchar(36) NOT NULL, + userid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + ts bigint NOT NULL, + linenum int NOT NULL, + linenumtemp boolean NOT NULL, + linetype varchar(10) NOT NULL, + linelocal boolean NOT NULL, + text text NOT NULL, + cmdid varchar(36) NOT NULL, + ephemeral boolean NOT NULL, + contentheight int NOT NULL, + star int NOT NULL, + archived boolean NOT NULL, + renderer varchar(50) NOT NULL, + bookmarked boolean NOT NULL, + PRIMARY KEY (screenid, lineid) +); +INSERT INTO new_line +SELECT screenid, userid, lineid, ts, linenum, linenumtemp, linetype, linelocal, + text, cmdid, ephemeral, contentheight, star, archived, renderer, bookmarked +FROM line; +DROP TABLE line; +ALTER TABLE new_line RENAME TO line; + +-- drop sessionid from cmd +CREATE TABLE new_cmd ( + screenid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + remotename varchar(50) NOT NULL, + cmdstr text NOT NULL, + rawcmdstr text NOT NULL, + festate json NOT NULL, + statebasehash varchar(36) NOT NULL, + statediffhasharr json NOT NULL, + termopts json NOT NULL, + origtermopts json NOT NULL, + status varchar(10) NOT NULL, + startpk json NOT NULL, + doneinfo json NOT NULL, + runout json NOT NULL, + rtnstate boolean NOT NULL, + rtnbasehash varchar(36) NOT NULL, + rtndiffhasharr json NOT NULL, + PRIMARY KEY (screenid, cmdid) +); +INSERT INTO new_cmd +SELECT screenid, cmdid, remoteownerid, remoteid, remotename, cmdstr, cmdstr, + festate, statebasehash, statediffhasharr, termopts, origtermopts, status, startpk, doneinfo, runout, rtnstate, rtnbasehash, rtndiffhasharr +FROM cmd; +DROP TABLE cmd; +ALTER TABLE new_cmd RENAME TO cmd; diff --git a/wavesrv/db/migrations/000014_simplifybookmarks.down.sql b/wavesrv/db/migrations/000014_simplifybookmarks.down.sql new file mode 100644 index 00000000..13119576 --- /dev/null +++ b/wavesrv/db/migrations/000014_simplifybookmarks.down.sql @@ -0,0 +1,9 @@ +CREATE TABLE IF NOT EXISTS "bookmark_cmd" ( + bookmarkid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + PRIMARY KEY (bookmarkid, screenid, cmdid) +); + +ALTER TABLE line ADD COLUMN bookmarked boolean NOT NULL DEFAULT 0; + diff --git a/wavesrv/db/migrations/000014_simplifybookmarks.up.sql b/wavesrv/db/migrations/000014_simplifybookmarks.up.sql new file mode 100644 index 00000000..fa2e43eb --- /dev/null +++ b/wavesrv/db/migrations/000014_simplifybookmarks.up.sql @@ -0,0 +1,3 @@ +DROP TABLE bookmark_cmd; +ALTER TABLE line DROP COLUMN bookmarked; + diff --git a/wavesrv/db/migrations/000015_lineupdates.down.sql b/wavesrv/db/migrations/000015_lineupdates.down.sql new file mode 100644 index 00000000..302819b2 --- /dev/null +++ b/wavesrv/db/migrations/000015_lineupdates.down.sql @@ -0,0 +1,4 @@ +DROP TABLE screenupdate; + +ALTER TABLE screen DROP COLUMN webshareopts; + diff --git a/wavesrv/db/migrations/000015_lineupdates.up.sql b/wavesrv/db/migrations/000015_lineupdates.up.sql new file mode 100644 index 00000000..93aad43f --- /dev/null +++ b/wavesrv/db/migrations/000015_lineupdates.up.sql @@ -0,0 +1,10 @@ +CREATE TABLE screenupdate ( + updateid integer PRIMARY KEY, + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + updatetype varchar(50) NOT NULL, + updatets bigint NOT NULL +); + +ALTER TABLE screen ADD COLUMN webshareopts json NOT NULL DEFAULT 'null'; + diff --git a/wavesrv/db/migrations/000016_webptypos.down.sql b/wavesrv/db/migrations/000016_webptypos.down.sql new file mode 100644 index 00000000..1d5c5243 --- /dev/null +++ b/wavesrv/db/migrations/000016_webptypos.down.sql @@ -0,0 +1,3 @@ +DROP TABLE webptypos; + +DROP INDEX idx_screenupdate_ids; diff --git a/wavesrv/db/migrations/000016_webptypos.up.sql b/wavesrv/db/migrations/000016_webptypos.up.sql new file mode 100644 index 00000000..e21618b0 --- /dev/null +++ b/wavesrv/db/migrations/000016_webptypos.up.sql @@ -0,0 +1,8 @@ +CREATE TABLE webptypos ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + ptypos bigint NOT NULL, + PRIMARY KEY (screenid, lineid) +); + +CREATE INDEX idx_screenupdate_ids ON screenupdate (screenid, lineid); diff --git a/wavesrv/db/migrations/000017_remotevars.down.sql b/wavesrv/db/migrations/000017_remotevars.down.sql new file mode 100644 index 00000000..442c1b8b --- /dev/null +++ b/wavesrv/db/migrations/000017_remotevars.down.sql @@ -0,0 +1,2 @@ +ALTER TABLE remote DROP COLUMN statevars; + diff --git a/wavesrv/db/migrations/000017_remotevars.up.sql b/wavesrv/db/migrations/000017_remotevars.up.sql new file mode 100644 index 00000000..f9136053 --- /dev/null +++ b/wavesrv/db/migrations/000017_remotevars.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE remote ADD COLUMN statevars json NOT NULL DEFAULT '{}'; + diff --git a/wavesrv/db/migrations/000018_modremote.down.sql b/wavesrv/db/migrations/000018_modremote.down.sql new file mode 100644 index 00000000..d3d73534 --- /dev/null +++ b/wavesrv/db/migrations/000018_modremote.down.sql @@ -0,0 +1,14 @@ +ALTER TABLE remote ADD COLUMN remotesudo; + +UPDATE remote +SET remotesudo = 1 +WHERE json_extract(sshopts, '$.issudo') +; + +UPDATE remote +SET sshopts = json_remove(sshopts, '$.issudo') +; + +ALTER TABLE remote ADD COLUMN physicalid varchar(36) NOT NULL DEFAULT ''; + +ALTER TABLE remote DROP COLUMN openaiopts; diff --git a/wavesrv/db/migrations/000018_modremote.up.sql b/wavesrv/db/migrations/000018_modremote.up.sql new file mode 100644 index 00000000..4cad763f --- /dev/null +++ b/wavesrv/db/migrations/000018_modremote.up.sql @@ -0,0 +1,11 @@ +UPDATE remote +SET sshopts = json_set(sshopts, '$.issudo', json('true')) +WHERE remotesudo +; + +ALTER TABLE remote DROP COLUMN remotesudo; + +ALTER TABLE remote DROP COLUMN physicalid; + +ALTER TABLE remote ADD COLUMN openaiopts json NOT NULL DEFAULT '{}'; + diff --git a/wavesrv/db/migrations/000019_clientopenai.down.sql b/wavesrv/db/migrations/000019_clientopenai.down.sql new file mode 100644 index 00000000..983f5d51 --- /dev/null +++ b/wavesrv/db/migrations/000019_clientopenai.down.sql @@ -0,0 +1 @@ +ALTER TABLE client DROP COLUMN openaiopts; diff --git a/wavesrv/db/migrations/000019_clientopenai.up.sql b/wavesrv/db/migrations/000019_clientopenai.up.sql new file mode 100644 index 00000000..a591cabc --- /dev/null +++ b/wavesrv/db/migrations/000019_clientopenai.up.sql @@ -0,0 +1,2 @@ +ALTER TABLE client ADD COLUMN openaiopts json NOT NULL DEFAULT '{}'; + diff --git a/wavesrv/db/migrations/000020_linecmd.down.sql b/wavesrv/db/migrations/000020_linecmd.down.sql new file mode 100644 index 00000000..6332dc5b --- /dev/null +++ b/wavesrv/db/migrations/000020_linecmd.down.sql @@ -0,0 +1,2 @@ +-- invalid, will throw an error, cannot migrate down +SELECT x; diff --git a/wavesrv/db/migrations/000020_linecmd.up.sql b/wavesrv/db/migrations/000020_linecmd.up.sql new file mode 100644 index 00000000..9e04bf69 --- /dev/null +++ b/wavesrv/db/migrations/000020_linecmd.up.sql @@ -0,0 +1,74 @@ +-- remove cmdid from line, history, and cmd (use lineid everywhere) + +CREATE TABLE cmd_new ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + remotename varchar(50) NOT NULL, + cmdstr text NOT NULL, + rawcmdstr text NOT NULL, + festate json NOT NULL, + statebasehash varchar(36) NOT NULL, + statediffhasharr json NOT NULL, + termopts json NOT NULL, + origtermopts json NOT NULL, + status varchar(10) NOT NULL, + cmdpid int NOT NULL, + remotepid int NOT NULL, + donets bigint NOT NULL, + exitcode int NOT NULL, + durationms int NOT NULL, + rtnstate boolean NOT NULL, + rtnbasehash varchar(36) NOT NULL, + rtndiffhasharr json NOT NULL, + runout json NOT NULL, + PRIMARY KEY (screenid, lineid) +); + +CREATE TABLE cmd_migrate20 ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + PRIMARY KEY (screenid, lineid) +); + +INSERT INTO cmd_migrate20 +SELECT screenid, lineid, cmdid +FROM line +WHERE cmdid <> '' +; + +INSERT INTO cmd_new +SELECT + c.screenid, + l.lineid, + c.remoteownerid, + c.remoteid, + c.remotename, + c.cmdstr, + c.rawcmdstr, + c.festate, + c.statebasehash, + c.statediffhasharr, + c.termopts, + c.origtermopts, + c.status, + coalesce(json_extract(startpk, '$.pid'), 0), + coalesce(json_extract(startpk, '$.mshellpid'), 0), + coalesce(json_extract(doneinfo, '$.ts'), 0), + coalesce(json_extract(doneinfo, '$.exitcode'), 0), + coalesce(json_extract(doneinfo, '$.durationms'), 0), + c.rtnstate, + c.rtnbasehash, + c.rtndiffhasharr, + c.runout +FROM cmd c +JOIN line l ON (l.cmdid = c.cmdid); + +DROP TABLE cmd; + +ALTER TABLE cmd_new RENAME TO cmd; + +ALTER TABLE history DROP COLUMN cmdid; +ALTER TABLE line DROP COLUMN cmdid; diff --git a/wavesrv/db/migrations/000021_linestate.down.sql b/wavesrv/db/migrations/000021_linestate.down.sql new file mode 100644 index 00000000..064dfbf6 --- /dev/null +++ b/wavesrv/db/migrations/000021_linestate.down.sql @@ -0,0 +1 @@ +ALTER TABLE line DROP COLUMN linestate; diff --git a/wavesrv/db/migrations/000021_linestate.up.sql b/wavesrv/db/migrations/000021_linestate.up.sql new file mode 100644 index 00000000..f53e0b51 --- /dev/null +++ b/wavesrv/db/migrations/000021_linestate.up.sql @@ -0,0 +1 @@ +ALTER TABLE line ADD COLUMN linestate json NOT NULL DEFAULT '{}'; diff --git a/wavesrv/db/migrations/000022_endwebshare.down.sql b/wavesrv/db/migrations/000022_endwebshare.down.sql new file mode 100644 index 00000000..b4eaa212 --- /dev/null +++ b/wavesrv/db/migrations/000022_endwebshare.down.sql @@ -0,0 +1 @@ +-- no down migration diff --git a/wavesrv/db/migrations/000022_endwebshare.up.sql b/wavesrv/db/migrations/000022_endwebshare.up.sql new file mode 100644 index 00000000..6e815154 --- /dev/null +++ b/wavesrv/db/migrations/000022_endwebshare.up.sql @@ -0,0 +1 @@ +UPDATE screen SET sharemode = 'local' AND webshareopts = 'null'; diff --git a/wavesrv/db/schema.sql b/wavesrv/db/schema.sql new file mode 100644 index 00000000..a2a8d9cd --- /dev/null +++ b/wavesrv/db/schema.sql @@ -0,0 +1,219 @@ +CREATE TABLE schema_migrations (version uint64,dirty bool); +CREATE UNIQUE INDEX version_unique ON schema_migrations (version); +CREATE TABLE client ( + clientid varchar(36) NOT NULL, + userid varchar(36) NOT NULL, + activesessionid varchar(36) NOT NULL, + userpublickeybytes blob NOT NULL, + userprivatekeybytes blob NOT NULL, + winsize json NOT NULL +, clientopts json NOT NULL DEFAULT '', feopts json NOT NULL DEFAULT '{}', cmdstoretype varchar(20) DEFAULT 'session', openaiopts json NOT NULL DEFAULT '{}'); +CREATE TABLE session ( + sessionid varchar(36) PRIMARY KEY, + name varchar(50) NOT NULL, + sessionidx int NOT NULL, + activescreenid varchar(36) NOT NULL, + notifynum int NOT NULL, + archived boolean NOT NULL, + archivedts bigint NOT NULL, + sharemode varchar(12) NOT NULL); +CREATE TABLE remote_instance ( + riid varchar(36) PRIMARY KEY, + name varchar(50) NOT NULL, + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + festate json NOT NULL, + statebasehash varchar(36) NOT NULL, + statediffhasharr json NOT NULL +); +CREATE TABLE state_base ( + basehash varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + version varchar(200) NOT NULL, + data blob NOT NULL +); +CREATE TABLE state_diff ( + diffhash varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + basehash varchar(36) NOT NULL, + diffhasharr json NOT NULL, + data blob NOT NULL +); +CREATE TABLE remote ( + remoteid varchar(36) PRIMARY KEY, + remotetype varchar(10) NOT NULL, + remotealias varchar(50) NOT NULL, + remotecanonicalname varchar(200) NOT NULL, + remoteuser varchar(50) NOT NULL, + remotehost varchar(200) NOT NULL, + connectmode varchar(20) NOT NULL, + autoinstall boolean NOT NULL, + sshopts json NOT NULL, + remoteopts json NOT NULL, + lastconnectts bigint NOT NULL, + local boolean NOT NULL, + archived boolean NOT NULL, + remoteidx int NOT NULL +, statevars json NOT NULL DEFAULT '{}', openaiopts json NOT NULL DEFAULT '{}'); +CREATE TABLE history ( + historyid varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + userid varchar(36) NOT NULL, + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + lineid int NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + remotename varchar(50) NOT NULL, + haderror boolean NOT NULL, + cmdstr text NOT NULL, + ismetacmd boolean, + incognito boolean +, linenum int NOT NULL DEFAULT 0); +CREATE TABLE activity ( + day varchar(20) PRIMARY KEY, + uploaded boolean NOT NULL, + tdata json NOT NULL, + tzname varchar(50) NOT NULL, + tzoffset int NOT NULL, + clientversion varchar(20) NOT NULL, + clientarch varchar(20) NOT NULL +, buildtime varchar(20) NOT NULL DEFAULT '-', osrelease varchar(20) NOT NULL DEFAULT '-'); +CREATE TABLE bookmark ( + bookmarkid varchar(36) PRIMARY KEY, + createdts bigint NOT NULL, + cmdstr text NOT NULL, + alias varchar(50) NOT NULL, + tags json NOT NULL, + description text NOT NULL +); +CREATE TABLE bookmark_order ( + tag varchar(50) NOT NULL, + bookmarkid varchar(36) NOT NULL, + orderidx int NOT NULL, + PRIMARY KEY (tag, bookmarkid) +); +CREATE TABLE playbook ( + playbookid varchar(36) PRIMARY KEY, + playbookname varchar(100) NOT NULL, + description text NOT NULL, + entryids json NOT NULL +); +CREATE TABLE playbook_entry ( + entryid varchar(36) PRIMARY KEY, + playbookid varchar(36) NOT NULL, + description text NOT NULL, + alias varchar(50) NOT NULL, + cmdstr text NOT NULL, + createdts bigint NOT NULL, + updatedts bigint NOT NULL +); +CREATE TABLE cloud_session ( + sessionid varchar(36) PRIMARY KEY, + viewkey varchar(50) NOT NULL, + writekey varchar(50) NOT NULL, + enckey varchar(100) NOT NULL, + enctype varchar(50) NOT NULL, + vts bigint NOT NULL, + acl json NOT NULL +); +CREATE TABLE cloud_update ( + updateid varchar(36) PRIMARY KEY, + ts bigint NOT NULL, + updatetype varchar(50) NOT NULL, + updatekeys json NOT NULL +); +CREATE TABLE cmd_migrate ( + sessionid varchar(36) NOT NULL, + screenid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL +); +CREATE TABLE IF NOT EXISTS "screen" ( + screenid varchar(36) NOT NULL, + sessionid varchar(36) NOT NULL, + name varchar(50) NOT NULL, + screenidx int NOT NULL, + screenopts json NOT NULL, + ownerid varchar(36) NOT NULL, + sharemode varchar(12) NOT NULL, + curremoteownerid varchar(36) NOT NULL, + curremoteid varchar(36) NOT NULL, + curremotename varchar(50) NOT NULL, + nextlinenum int NOT NULL, + selectedline int NOT NULL, + anchor json NOT NULL, + focustype varchar(12) NOT NULL, + archived boolean NOT NULL, + archivedts bigint NOT NULL, webshareopts json NOT NULL DEFAULT 'null', + PRIMARY KEY (screenid) +); +CREATE TABLE IF NOT EXISTS "line" ( + screenid varchar(36) NOT NULL, + userid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + ts bigint NOT NULL, + linenum int NOT NULL, + linenumtemp boolean NOT NULL, + linetype varchar(10) NOT NULL, + linelocal boolean NOT NULL, + text text NOT NULL, + ephemeral boolean NOT NULL, + contentheight int NOT NULL, + star int NOT NULL, + archived boolean NOT NULL, + renderer varchar(50) NOT NULL, linestate json NOT NULL DEFAULT '{}', + PRIMARY KEY (screenid, lineid) +); +CREATE TABLE screenupdate ( + updateid integer PRIMARY KEY, + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + updatetype varchar(50) NOT NULL, + updatets bigint NOT NULL +); +CREATE TABLE webptypos ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + ptypos bigint NOT NULL, + PRIMARY KEY (screenid, lineid) +); +CREATE INDEX idx_screenupdate_ids ON screenupdate (screenid, lineid); +CREATE TABLE cmd_migration ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + PRIMARY KEY (screenid, lineid) +); +CREATE TABLE IF NOT EXISTS "cmd" ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + remoteownerid varchar(36) NOT NULL, + remoteid varchar(36) NOT NULL, + remotename varchar(50) NOT NULL, + cmdstr text NOT NULL, + rawcmdstr text NOT NULL, + festate json NOT NULL, + statebasehash varchar(36) NOT NULL, + statediffhasharr json NOT NULL, + termopts json NOT NULL, + origtermopts json NOT NULL, + status varchar(10) NOT NULL, + cmdpid int NOT NULL, + remotepid int NOT NULL, + donets bigint NOT NULL, + exitcode int NOT NULL, + durationms int NOT NULL, + rtnstate boolean NOT NULL, + rtnbasehash varchar(36) NOT NULL, + rtndiffhasharr json NOT NULL, + runout json NOT NULL, + PRIMARY KEY (screenid, lineid) +); +CREATE TABLE cmd_migrate20 ( + screenid varchar(36) NOT NULL, + lineid varchar(36) NOT NULL, + cmdid varchar(36) NOT NULL, + PRIMARY KEY (screenid, lineid) +); diff --git a/wavesrv/go.mod b/wavesrv/go.mod new file mode 100644 index 00000000..a51314ea --- /dev/null +++ b/wavesrv/go.mod @@ -0,0 +1,28 @@ +module github.com/commandlinedev/prompt-server + +go 1.18 + +require ( + github.com/alessio/shellescape v1.4.1 + github.com/armon/circbuf v0.0.0-20190214190532-5111143e8da2 + github.com/commandlinedev/apishell v0.0.0 + github.com/creack/pty v1.1.18 + github.com/golang-migrate/migrate/v4 v4.16.2 + github.com/google/uuid v1.3.0 + github.com/gorilla/mux v1.8.0 + github.com/gorilla/websocket v1.5.0 + github.com/jmoiron/sqlx v1.3.5 + github.com/mattn/go-sqlite3 v1.14.16 + github.com/sashabaranov/go-openai v1.9.0 + github.com/sawka/txwrap v0.1.2 + golang.org/x/crypto v0.7.0 + golang.org/x/mod v0.10.0 + golang.org/x/sys v0.10.0 + mvdan.cc/sh/v3 v3.7.0 +) + +require ( + github.com/hashicorp/errwrap v1.1.0 // indirect + github.com/hashicorp/go-multierror v1.1.1 // indirect + go.uber.org/atomic v1.7.0 // indirect +) \ No newline at end of file diff --git a/wavesrv/go.sum b/wavesrv/go.sum new file mode 100644 index 00000000..7c819d2f --- /dev/null +++ b/wavesrv/go.sum @@ -0,0 +1,64 @@ +github.com/alessio/shellescape v1.4.1 h1:V7yhSDDn8LP4lc4jS8pFkt0zCnzVJlG5JXy9BVKJUX0= +github.com/alessio/shellescape v1.4.1/go.mod h1:PZAiSCk0LJaZkiCSkPv8qIobYglO3FPpyFjDCtHLS30= +github.com/armon/circbuf v0.0.0-20190214190532-5111143e8da2 h1:7Ip0wMmLHLRJdrloDxZfhMm0xrLXZS8+COSu2bXmEQs= +github.com/armon/circbuf v0.0.0-20190214190532-5111143e8da2/go.mod h1:3U/XgcO3hCbHZ8TKRvWD2dDTCfh9M9ya+I9JpbB7O8o= +github.com/creack/pty v1.1.18 h1:n56/Zwd5o6whRC5PMGretI4IdRLlmBXYNjScPaBgsbY= +github.com/creack/pty v1.1.18/go.mod h1:MOBLtS5ELjhRRrroQr9kyvTxUAFNvYEK993ew/Vr4O4= +github.com/davecgh/go-spew v1.1.0/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/davecgh/go-spew v1.1.1 h1:vj9j/u1bqnvCEfJOwUhtlOARqs3+rkHYY13jYWTU97c= +github.com/davecgh/go-spew v1.1.1/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38= +github.com/frankban/quicktest v1.14.5 h1:dfYrrRyLtiqT9GyKXgdh+k4inNeTvmGbuSgZ3lx3GhA= +github.com/frankban/quicktest v1.14.5/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7zb5vbUoiM6w0= +github.com/go-sql-driver/mysql v1.6.0 h1:BCTh4TKNUYmOmMUcQ3IipzF5prigylS7XXjEkfCHuOE= +github.com/go-sql-driver/mysql v1.6.0/go.mod h1:DCzpHaOWr8IXmIStZouvnhqoel9Qv2LBy8hT2VhHyBg= +github.com/golang-migrate/migrate/v4 v4.16.2 h1:8coYbMKUyInrFk1lfGfRovTLAW7PhWp8qQDT2iKfuoA= +github.com/golang-migrate/migrate/v4 v4.16.2/go.mod h1:pfcJX4nPHaVdc5nmdCikFBWtm+UBpiZjRNNsyBbp0/o= +github.com/google/go-cmp v0.5.9 h1:O2Tfq5qg4qc4AmwVlvv0oLiVAGB7enBSJ2x2DqQFi38= +github.com/google/go-cmp v0.5.9/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/uuid v1.3.0 h1:t6JiXgmwXMjEs8VusXIJk2BXHsn+wx8BZdTaoZ5fu7I= +github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= +github.com/gorilla/mux v1.8.0 h1:i40aqfkR1h2SlN9hojwV5ZA91wcXFOvkdNIeFDP5koI= +github.com/gorilla/mux v1.8.0/go.mod h1:DVbg23sWSpFRCP0SfiEN6jmj59UnW/n46BH5rLB71So= +github.com/gorilla/websocket v1.5.0 h1:PPwGk2jz7EePpoHN/+ClbZu8SPxiqlu12wZP/3sWmnc= +github.com/gorilla/websocket v1.5.0/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= +github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= +github.com/hashicorp/errwrap v1.1.0 h1:OxrOeh75EUXMY8TBjag2fzXGZ40LB6IKw45YeGUDY2I= +github.com/hashicorp/errwrap v1.1.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= +github.com/hashicorp/go-multierror v1.1.1 h1:H5DkEtf6CXdFp0N0Em5UCwQpXMWke8IA0+lD48awMYo= +github.com/hashicorp/go-multierror v1.1.1/go.mod h1:iw975J/qwKPdAO1clOe2L8331t/9/fmwbPZ6JB6eMoM= +github.com/jmoiron/sqlx v1.3.5 h1:vFFPA71p1o5gAeqtEAwLU4dnX2napprKtHr7PYIcN3g= +github.com/jmoiron/sqlx v1.3.5/go.mod h1:nRVWtLre0KfCLJvgxzCsLVMogSvQ1zNJtpYr2Ccp0mQ= +github.com/kr/pretty v0.3.1 h1:flRD4NNwYAUpkphVc1HcthR4KEIFJ65n8Mw5qdRn3LE= +github.com/kr/pretty v0.3.1/go.mod h1:hoEshYVHaxMs3cyo3Yncou5ZscifuDolrwPKZanG3xk= +github.com/kr/text v0.2.0 h1:5Nx0Ya0ZqY2ygV366QzturHI13Jq95ApcVaJBhpS+AY= +github.com/kr/text v0.2.0/go.mod h1:eLer722TekiGuMkidMxC/pM04lWEeraHUUmBw8l2grE= +github.com/lib/pq v1.2.0/go.mod h1:5WUZQaWbwv1U+lTReE5YruASi9Al49XbQIvNi/34Woo= +github.com/lib/pq v1.10.2 h1:AqzbZs4ZoCBp+GtejcpCpcxM3zlSMx29dXbUSeVtJb8= +github.com/lib/pq v1.10.2/go.mod h1:AlVN5x4E4T544tWzH6hKfbfQvm3HdbOxrmggDNAPY9o= +github.com/mattn/go-sqlite3 v1.14.6/go.mod h1:NyWgC/yNuGj7Q9rpYnZvas74GogHl5/Z4A/KQRfk6bU= +github.com/mattn/go-sqlite3 v1.14.16 h1:yOQRA0RpS5PFz/oikGwBEqvAWhWg5ufRz4ETLjwpU1Y= +github.com/mattn/go-sqlite3 v1.14.16/go.mod h1:2eHXhiwb8IkHr+BDWZGa96P6+rkvnG63S2DGjv9HUNg= +github.com/pmezard/go-difflib v1.0.0 h1:4DBwDE0NGyQoBHbLQYPwSUPoCMWR5BEzIk/f1lZbAQM= +github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= +github.com/rogpeppe/go-internal v1.10.1-0.20230524175051-ec119421bb97 h1:3RPlVWzZ/PDqmVuf/FKHARG5EMid/tl7cv54Sw/QRVY= +github.com/rogpeppe/go-internal v1.10.1-0.20230524175051-ec119421bb97/go.mod h1:ddIwULY96R17DhadqLgMfk9H9tvdUzkipdSkR5nkCZA= +github.com/sashabaranov/go-openai v1.9.0 h1:NoiO++IISxxJ1pRc0n7uZvMGMake0G+FJ1XPwXtprsA= +github.com/sashabaranov/go-openai v1.9.0/go.mod h1:lj5b/K+zjTSFxVLijLSTDZuP7adOgerWeFyZLUhAKRg= +github.com/sawka/txwrap v0.1.2 h1:v8xS0Z1LE7/6vMZA81PYihI+0TSR6Zm1MalzzBIuXKc= +github.com/sawka/txwrap v0.1.2/go.mod h1:T3nlw2gVpuolo6/XEetvBbk1oMXnY978YmBFy1UyHvw= +github.com/stretchr/objx v0.1.0/go.mod h1:HFkY916IF+rwdDfMAkV7OtwuqBVzrE8GR6GFx+wExME= +github.com/stretchr/testify v1.3.0/go.mod h1:M5WIy9Dh21IEIfnGCwXGc5bZfKNJtfHm1UVUgZn+9EI= +github.com/stretchr/testify v1.8.1 h1:w7B6lhMri9wdJUVmEZPGGhZzrYTPvgJArz7wNPgYKsk= +github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= +go.uber.org/atomic v1.7.0 h1:ADUqmZGgLDDfbSL9ZmPxKTybcoEYHgpYfELNoN+7hsw= +go.uber.org/atomic v1.7.0/go.mod h1:fEN4uk6kAWBTFdckzkM89CLk9XfWZrxpCo0nPH17wJc= +golang.org/x/crypto v0.7.0 h1:AvwMYaRytfdeVt3u6mLaxYtErKYjxA2OXjJ1HHq6t3A= +golang.org/x/crypto v0.7.0/go.mod h1:pYwdfH91IfpZVANVyUOhSIPZaFoJGxTFbZhFTx+dXZU= +golang.org/x/mod v0.10.0 h1:lFO9qtOdlre5W1jxS3r/4szv2/6iXxScdzjoBMXNhYk= +golang.org/x/mod v0.10.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= +golang.org/x/sys v0.10.0 h1:SqMFp9UcQJZa+pmYuAKjd9xq1f0j5rLcDIk0mj4qAsA= +golang.org/x/sys v0.10.0/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= +gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= +gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= +mvdan.cc/sh/v3 v3.7.0 h1:lSTjdP/1xsddtaKfGg7Myu7DnlHItd3/M2tomOcNNBg= +mvdan.cc/sh/v3 v3.7.0/go.mod h1:K2gwkaesF/D7av7Kxl0HbF5kGOd2ArupNTX3X44+8l8= diff --git a/wavesrv/pkg/cmdrunner/cmdrunner.go b/wavesrv/pkg/cmdrunner/cmdrunner.go new file mode 100644 index 00000000..10f5f99e --- /dev/null +++ b/wavesrv/pkg/cmdrunner/cmdrunner.go @@ -0,0 +1,3828 @@ +package cmdrunner + +import ( + "bytes" + "context" + "encoding/json" + "fmt" + "io/fs" + "log" + "net/url" + "os" + "path/filepath" + "regexp" + "sort" + "strconv" + "strings" + "syscall" + "time" + "unicode" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/prompt-server/pkg/comp" + "github.com/commandlinedev/prompt-server/pkg/dbutil" + "github.com/commandlinedev/prompt-server/pkg/pcloud" + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/remote/openai" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/commandlinedev/prompt-server/pkg/scpacket" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/commandlinedev/prompt-server/pkg/utilfn" + "github.com/google/uuid" +) + +const ( + HistoryTypeScreen = "screen" + HistoryTypeSession = "session" + HistoryTypeGlobal = "global" +) + +func init() { + comp.RegisterSimpleCompFn(comp.CGTypeMeta, simpleCompMeta) + comp.RegisterSimpleCompFn(comp.CGTypeCommandMeta, simpleCompCommandMeta) +} + +const DefaultUserId = "user" +const MaxNameLen = 50 +const MaxShareNameLen = 150 +const MaxRendererLen = 50 +const MaxRemoteAliasLen = 50 +const PasswordUnchangedSentinel = "--unchanged--" +const DefaultPTERM = "MxM" +const MaxCommandLen = 4096 +const MaxSignalLen = 12 +const MaxSignalNum = 64 +const MaxEvalDepth = 5 +const MaxOpenAIAPITokenLen = 100 +const MaxOpenAIModelLen = 100 + +const TsFormatStr = "2006-01-02 15:04:05" + +const ( + KwArgRenderer = "renderer" + KwArgView = "view" + KwArgState = "state" + KwArgTemplate = "template" + KwArgLang = "lang" +) + +var ColorNames = []string{"black", "red", "green", "yellow", "blue", "magenta", "cyan", "white", "orange"} +var RemoteColorNames = []string{"red", "green", "yellow", "blue", "magenta", "cyan", "white", "orange"} +var RemoteSetArgs = []string{"alias", "connectmode", "key", "password", "autoinstall", "color"} + +var ScreenCmds = []string{"run", "comment", "cd", "cr", "clear", "sw", "reset", "signal", "chat"} +var NoHistCmds = []string{"_compgen", "line", "history", "_killserver"} +var GlobalCmds = []string{"session", "screen", "remote", "set", "client", "telemetry", "bookmark", "bookmarks"} + +var SetVarNameMap map[string]string = map[string]string{ + "tabcolor": "screen.tabcolor", + "pterm": "screen.pterm", + "anchor": "screen.anchor", + "focus": "screen.focus", + "line": "screen.line", +} + +var SetVarScopes = []SetVarScope{ + SetVarScope{ScopeName: "global", VarNames: []string{}}, + SetVarScope{ScopeName: "client", VarNames: []string{"telemetry"}}, + SetVarScope{ScopeName: "session", VarNames: []string{"name", "pos"}}, + SetVarScope{ScopeName: "screen", VarNames: []string{"name", "tabcolor", "pos", "pterm", "anchor", "focus", "line"}}, + SetVarScope{ScopeName: "line", VarNames: []string{}}, + // connection = remote, remote = remoteinstance + SetVarScope{ScopeName: "connection", VarNames: []string{"alias", "connectmode", "key", "password", "autoinstall", "color"}}, + SetVarScope{ScopeName: "remote", VarNames: []string{}}, +} + +var hostNameRe = regexp.MustCompile("^[a-z][a-z0-9.-]*$") +var userHostRe = regexp.MustCompile("^(sudo@)?([a-z][a-z0-9-]*)@([a-z0-9][a-z0-9.-]*)(?::([0-9]+))?$") +var remoteAliasRe = regexp.MustCompile("^[a-zA-Z][a-zA-Z0-9_-]*$") +var genericNameRe = regexp.MustCompile("^[a-zA-Z][a-zA-Z0-9_ .()<>,/\"'\\[\\]{}=+$@!*-]*$") +var rendererRe = regexp.MustCompile("^[a-zA-Z][a-zA-Z0-9_.:-]*$") +var positionRe = regexp.MustCompile("^((S?\\+|E?-)?[0-9]+|(\\+|-|S|E))$") +var wsRe = regexp.MustCompile("\\s+") +var sigNameRe = regexp.MustCompile("^((SIG[A-Z0-9]+)|(\\d+))$") + +type contextType string + +var historyContextKey = contextType("history") +var depthContextKey = contextType("depth") + +type SetVarScope struct { + ScopeName string + VarNames []string +} + +type historyContextType struct { + LineId string + LineNum int64 + RemotePtr *sstore.RemotePtrType +} + +type MetaCmdFnType = func(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) +type MetaCmdEntryType struct { + IsAlias bool + Fn MetaCmdFnType +} + +var MetaCmdFnMap = make(map[string]MetaCmdEntryType) + +func init() { + registerCmdFn("run", RunCommand) + registerCmdFn("eval", EvalCommand) + registerCmdFn("comment", CommentCommand) + registerCmdFn("cr", CrCommand) + registerCmdFn("connect", CrCommand) + registerCmdFn("_compgen", CompGenCommand) + registerCmdFn("clear", ClearCommand) + registerCmdFn("reset", RemoteResetCommand) + registerCmdFn("signal", SignalCommand) + registerCmdFn("sync", SyncCommand) + + registerCmdFn("session", SessionCommand) + registerCmdFn("session:open", SessionOpenCommand) + registerCmdAlias("session:new", SessionOpenCommand) + registerCmdFn("session:set", SessionSetCommand) + registerCmdAlias("session:delete", SessionDeleteCommand) + registerCmdFn("session:purge", SessionDeleteCommand) + registerCmdFn("session:archive", SessionArchiveCommand) + registerCmdFn("session:showall", SessionShowAllCommand) + registerCmdFn("session:show", SessionShowCommand) + registerCmdFn("session:openshared", SessionOpenSharedCommand) + + registerCmdFn("screen", ScreenCommand) + registerCmdFn("screen:archive", ScreenArchiveCommand) + registerCmdFn("screen:purge", ScreenPurgeCommand) + registerCmdFn("screen:open", ScreenOpenCommand) + registerCmdAlias("screen:new", ScreenOpenCommand) + registerCmdFn("screen:set", ScreenSetCommand) + registerCmdFn("screen:showall", ScreenShowAllCommand) + registerCmdFn("screen:reset", ScreenResetCommand) + registerCmdFn("screen:webshare", ScreenWebShareCommand) + + registerCmdAlias("remote", RemoteCommand) + registerCmdFn("remote:show", RemoteShowCommand) + registerCmdFn("remote:showall", RemoteShowAllCommand) + registerCmdFn("remote:new", RemoteNewCommand) + registerCmdFn("remote:archive", RemoteArchiveCommand) + registerCmdFn("remote:set", RemoteSetCommand) + registerCmdFn("remote:disconnect", RemoteDisconnectCommand) + registerCmdFn("remote:connect", RemoteConnectCommand) + registerCmdFn("remote:install", RemoteInstallCommand) + registerCmdFn("remote:installcancel", RemoteInstallCancelCommand) + registerCmdFn("remote:reset", RemoteResetCommand) + + registerCmdFn("screen:resize", ScreenResizeCommand) + + registerCmdFn("line", LineCommand) + registerCmdFn("line:show", LineShowCommand) + registerCmdFn("line:star", LineStarCommand) + registerCmdFn("line:bookmark", LineBookmarkCommand) + registerCmdFn("line:pin", LinePinCommand) + registerCmdFn("line:archive", LineArchiveCommand) + registerCmdFn("line:purge", LinePurgeCommand) + registerCmdFn("line:setheight", LineSetHeightCommand) + registerCmdFn("line:view", LineViewCommand) + registerCmdFn("line:set", LineSetCommand) + + registerCmdFn("client", ClientCommand) + registerCmdFn("client:show", ClientShowCommand) + registerCmdFn("client:set", ClientSetCommand) + registerCmdFn("client:notifyupdatewriter", ClientNotifyUpdateWriterCommand) + registerCmdFn("client:accepttos", ClientAcceptTosCommand) + + registerCmdFn("telemetry", TelemetryCommand) + registerCmdFn("telemetry:on", TelemetryOnCommand) + registerCmdFn("telemetry:off", TelemetryOffCommand) + registerCmdFn("telemetry:send", TelemetrySendCommand) + registerCmdFn("telemetry:show", TelemetryShowCommand) + + registerCmdFn("history", HistoryCommand) + registerCmdFn("history:viewall", HistoryViewAllCommand) + registerCmdFn("history:purge", HistoryPurgeCommand) + + registerCmdFn("bookmarks:show", BookmarksShowCommand) + + registerCmdFn("bookmark:set", BookmarkSetCommand) + registerCmdFn("bookmark:delete", BookmarkDeleteCommand) + + registerCmdFn("chat", OpenAICommand) + + registerCmdFn("_killserver", KillServerCommand) + + registerCmdFn("set", SetCommand) + + registerCmdFn("view:stat", ViewStatCommand) + registerCmdFn("view:test", ViewTestCommand) + + registerCmdFn("edit:test", EditTestCommand) + + // CodeEditCommand is overloaded to do codeedit and codeview + registerCmdFn("codeedit", CodeEditCommand) + registerCmdFn("codeview", CodeEditCommand) + + registerCmdFn("imageview", ImageViewCommand) + registerCmdFn("mdview", MarkdownViewCommand) + registerCmdFn("markdownview", MarkdownViewCommand) + + registerCmdFn("csvview", CSVViewCommand) +} + +func getValidCommands() []string { + var rtn []string + for key, val := range MetaCmdFnMap { + if val.IsAlias { + continue + } + rtn = append(rtn, "/"+key) + } + return rtn +} + +func registerCmdFn(cmdName string, fn MetaCmdFnType) { + MetaCmdFnMap[cmdName] = MetaCmdEntryType{Fn: fn} +} + +func registerCmdAlias(cmdName string, fn MetaCmdFnType) { + MetaCmdFnMap[cmdName] = MetaCmdEntryType{IsAlias: true, Fn: fn} +} + +func GetCmdStr(pk *scpacket.FeCommandPacketType) string { + if pk.MetaSubCmd == "" { + return pk.MetaCmd + } + return pk.MetaCmd + ":" + pk.MetaSubCmd +} + +func HandleCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + metaCmd := SubMetaCmd(pk.MetaCmd) + var cmdName string + if pk.MetaSubCmd == "" { + cmdName = metaCmd + } else { + cmdName = fmt.Sprintf("%s:%s", pk.MetaCmd, pk.MetaSubCmd) + } + entry := MetaCmdFnMap[cmdName] + if entry.Fn == nil { + if MetaCmdFnMap[metaCmd].Fn != nil { + return nil, fmt.Errorf("invalid /%s subcommand '%s'", metaCmd, pk.MetaSubCmd) + } + return nil, fmt.Errorf("invalid command '/%s', no handler", cmdName) + } + return entry.Fn(ctx, pk) +} + +func firstArg(pk *scpacket.FeCommandPacketType) string { + if len(pk.Args) == 0 { + return "" + } + return pk.Args[0] +} + +func argN(pk *scpacket.FeCommandPacketType, n int) string { + if len(pk.Args) <= n { + return "" + } + return pk.Args[n] +} + +func resolveBool(arg string, def bool) bool { + if arg == "" { + return def + } + if arg == "0" || arg == "false" { + return false + } + return true +} + +func defaultStr(arg string, def string) string { + if arg == "" { + return def + } + return arg +} + +func resolveFile(arg string) (string, error) { + if arg == "" { + return "", nil + } + fileName := base.ExpandHomeDir(arg) + if !strings.HasPrefix(fileName, "/") { + return "", fmt.Errorf("must be absolute, cannot be a relative path") + } + fd, err := os.Open(fileName) + if fd != nil { + fd.Close() + } + if err != nil { + return "", fmt.Errorf("cannot open file: %v", err) + } + return fileName, nil +} + +func resolvePosInt(arg string, def int) (int, error) { + if arg == "" { + return def, nil + } + ival, err := strconv.Atoi(arg) + if err != nil { + return 0, err + } + if ival <= 0 { + return 0, fmt.Errorf("must be greater than 0") + } + return ival, nil +} + +func isAllDigits(arg string) bool { + if len(arg) == 0 { + return false + } + for i := 0; i < len(arg); i++ { + if arg[i] >= '0' && arg[i] <= '9' { + continue + } + return false + } + return true +} + +func resolveNonNegInt(arg string, def int) (int, error) { + if arg == "" { + return def, nil + } + ival, err := strconv.Atoi(arg) + if err != nil { + return 0, err + } + if ival < 0 { + return 0, fmt.Errorf("cannot be negative") + } + return ival, nil +} + +var histExpansionRe = regexp.MustCompile(`^!(\d+)$`) + +func doCmdHistoryExpansion(ctx context.Context, ids resolvedIds, cmdStr string) (string, error) { + if !strings.HasPrefix(cmdStr, "!") { + return "", nil + } + if strings.HasPrefix(cmdStr, "! ") { + return "", nil + } + if cmdStr == "!!" { + return doHistoryExpansion(ctx, ids, -1) + } + if strings.HasPrefix(cmdStr, "!-") { + return "", fmt.Errorf("prompt does not support negative history offsets, use a stable positive history offset instead: '![linenum]'") + } + m := histExpansionRe.FindStringSubmatch(cmdStr) + if m == nil { + return "", fmt.Errorf("unsupported history substitution, can use '!!' or '![linenum]'") + } + ival, err := strconv.Atoi(m[1]) + if err != nil { + return "", fmt.Errorf("invalid history expansion") + } + return doHistoryExpansion(ctx, ids, ival) +} + +func doHistoryExpansion(ctx context.Context, ids resolvedIds, hnum int) (string, error) { + if hnum == 0 { + return "", fmt.Errorf("invalid history expansion, cannot expand line number '0'") + } + if hnum < -1 { + return "", fmt.Errorf("invalid history expansion, cannot expand negative history offsets") + } + foundHistoryNum := hnum + if hnum == -1 { + var err error + foundHistoryNum, err = sstore.GetLastHistoryLineNum(ctx, ids.ScreenId) + if err != nil { + return "", fmt.Errorf("cannot expand history, error finding last history item: %v", err) + } + if foundHistoryNum == 0 { + return "", fmt.Errorf("cannot expand history, no last history item") + } + } + hitem, err := sstore.GetHistoryItemByLineNum(ctx, ids.ScreenId, foundHistoryNum) + if err != nil { + return "", fmt.Errorf("cannot get history item '%d': %v", foundHistoryNum, err) + } + if hitem == nil { + return "", fmt.Errorf("cannot expand history, history item '%d' not found", foundHistoryNum) + } + return hitem.CmdStr, nil +} + +func getEvalDepth(ctx context.Context) int { + depthVal := ctx.Value(depthContextKey) + if depthVal == nil { + return 0 + } + return depthVal.(int) +} + +func SyncCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, fmt.Errorf("/run error: %w", err) + } + runPacket := packet.MakeRunPacket() + runPacket.ReqId = uuid.New().String() + runPacket.CK = base.MakeCommandKey(ids.ScreenId, scbase.GenPromptUUID()) + runPacket.UsePty = true + ptermVal := defaultStr(pk.Kwargs["pterm"], DefaultPTERM) + runPacket.TermOpts, err = GetUITermOpts(pk.UIContext.WinSize, ptermVal) + if err != nil { + return nil, fmt.Errorf("/sync error, invalid 'pterm' value %q: %v", ptermVal, err) + } + runPacket.Command = ":" + runPacket.ReturnState = true + cmd, callback, err := remote.RunCommand(ctx, ids.SessionId, ids.ScreenId, ids.Remote.RemotePtr, runPacket) + if callback != nil { + defer callback() + } + if err != nil { + return nil, err + } + cmd.RawCmdStr = pk.GetRawStr() + update, err := addLineForCmd(ctx, "/sync", true, ids, cmd, "terminal", nil) + if err != nil { + return nil, err + } + update.Interactive = pk.Interactive + sstore.MainBus.SendScreenUpdate(ids.ScreenId, update) + return nil, nil +} + +func getRendererArg(pk *scpacket.FeCommandPacketType) (string, error) { + rval := pk.Kwargs[KwArgView] + if rval == "" { + rval = pk.Kwargs[KwArgRenderer] + } + if rval == "" { + return "", nil + } + err := validateRenderer(rval) + if err != nil { + return "", err + } + return rval, nil +} + +func getTemplateArg(pk *scpacket.FeCommandPacketType) (string, error) { + rval := pk.Kwargs[KwArgTemplate] + if rval == "" { + return "", nil + } + // TODO validate + return rval, nil +} + +func getLangArg(pk *scpacket.FeCommandPacketType) (string, error) { + // TODO better error checking + if len(pk.Kwargs[KwArgLang]) > 50 { + return "", nil // TODO return error, don't fail silently + } + return pk.Kwargs[KwArgLang], nil +} + +func RunCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, fmt.Errorf("/run error: %w", err) + } + renderer, err := getRendererArg(pk) + if err != nil { + return nil, fmt.Errorf("/run error, invalid view/renderer: %w", err) + } + templateArg, err := getTemplateArg(pk) + if err != nil { + return nil, fmt.Errorf("/run error, invalid template: %w", err) + } + langArg, err := getLangArg(pk) + if err != nil { + return nil, fmt.Errorf("/run error, invalid lang: %w", err) + } + cmdStr := firstArg(pk) + expandedCmdStr, err := doCmdHistoryExpansion(ctx, ids, cmdStr) + if err != nil { + return nil, err + } + if expandedCmdStr != "" { + newPk := scpacket.MakeFeCommandPacket() + newPk.MetaCmd = "eval" + newPk.Args = []string{expandedCmdStr} + newPk.Kwargs = pk.Kwargs + newPk.RawStr = pk.RawStr + newPk.UIContext = pk.UIContext + newPk.Interactive = pk.Interactive + evalDepth := getEvalDepth(ctx) + ctxWithDepth := context.WithValue(ctx, depthContextKey, evalDepth+1) + return EvalCommand(ctxWithDepth, newPk) + } + isRtnStateCmd := IsReturnStateCommand(cmdStr) + // runPacket.State is set in remote.RunCommand() + runPacket := packet.MakeRunPacket() + runPacket.ReqId = uuid.New().String() + runPacket.CK = base.MakeCommandKey(ids.ScreenId, scbase.GenPromptUUID()) + runPacket.UsePty = true + ptermVal := defaultStr(pk.Kwargs["pterm"], DefaultPTERM) + runPacket.TermOpts, err = GetUITermOpts(pk.UIContext.WinSize, ptermVal) + if err != nil { + return nil, fmt.Errorf("/run error, invalid 'pterm' value %q: %v", ptermVal, err) + } + runPacket.Command = strings.TrimSpace(cmdStr) + runPacket.ReturnState = resolveBool(pk.Kwargs["rtnstate"], isRtnStateCmd) + cmd, callback, err := remote.RunCommand(ctx, ids.SessionId, ids.ScreenId, ids.Remote.RemotePtr, runPacket) + if callback != nil { + defer callback() + } + if err != nil { + return nil, err + } + cmd.RawCmdStr = pk.GetRawStr() + lineState := make(map[string]any) + if templateArg != "" { + lineState[sstore.LineState_Template] = templateArg + } + if langArg != "" { + lineState[sstore.LineState_Lang] = langArg + } + update, err := addLineForCmd(ctx, "/run", true, ids, cmd, renderer, lineState) + if err != nil { + return nil, err + } + update.Interactive = pk.Interactive + sstore.MainBus.SendScreenUpdate(ids.ScreenId, update) + return nil, nil +} + +func addToHistory(ctx context.Context, pk *scpacket.FeCommandPacketType, historyContext historyContextType, isMetaCmd bool, hadError bool) error { + cmdStr := firstArg(pk) + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return err + } + isIncognito, err := sstore.IsIncognitoScreen(ctx, ids.SessionId, ids.ScreenId) + if err != nil { + return fmt.Errorf("cannot add to history, error looking up incognito status of screen: %v", err) + } + hitem := &sstore.HistoryItemType{ + HistoryId: scbase.GenPromptUUID(), + Ts: time.Now().UnixMilli(), + UserId: DefaultUserId, + SessionId: ids.SessionId, + ScreenId: ids.ScreenId, + LineId: historyContext.LineId, + LineNum: historyContext.LineNum, + HadError: hadError, + CmdStr: cmdStr, + IsMetaCmd: isMetaCmd, + Incognito: isIncognito, + } + if !isMetaCmd && historyContext.RemotePtr != nil { + hitem.Remote = *historyContext.RemotePtr + } + err = sstore.InsertHistoryItem(ctx, hitem) + if err != nil { + return err + } + return nil +} + +func EvalCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("usage: /eval [command], no command passed to eval") + } + if len(pk.Args[0]) > MaxCommandLen { + return nil, fmt.Errorf("command length too long len:%d, max:%d", len(pk.Args[0]), MaxCommandLen) + } + evalDepth := getEvalDepth(ctx) + if pk.Interactive && evalDepth == 0 { + err := sstore.UpdateCurrentActivity(ctx, sstore.ActivityUpdate{NumCommands: 1}) + if err != nil { + log.Printf("[error] incrementing activity numcommands: %v\n", err) + // fall through (non-fatal error) + } + } + if evalDepth > MaxEvalDepth { + return nil, fmt.Errorf("alias/history expansion max-depth exceeded") + } + var historyContext historyContextType + ctxWithHistory := context.WithValue(ctx, historyContextKey, &historyContext) + var update sstore.UpdatePacket + newPk, rtnErr := EvalMetaCommand(ctxWithHistory, pk) + if rtnErr == nil { + update, rtnErr = HandleCommand(ctxWithHistory, newPk) + } + if !resolveBool(pk.Kwargs["nohist"], false) { + err := addToHistory(ctx, pk, historyContext, (newPk.MetaCmd != "run"), (rtnErr != nil)) + if err != nil { + log.Printf("[error] adding to history: %v\n", err) + // fall through (non-fatal error) + } + } + return update, rtnErr +} + +func ScreenArchiveCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) // don't force R_Screen + if err != nil { + return nil, fmt.Errorf("/screen:archive cannot archive screen: %w", err) + } + screenId := ids.ScreenId + if len(pk.Args) > 0 { + ri, err := resolveSessionScreen(ctx, ids.SessionId, pk.Args[0], ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("/screen:archive cannot resolve screen arg: %v", err) + } + screenId = ri.Id + } + if screenId == "" { + return nil, fmt.Errorf("/screen:archive no active screen or screen arg passed") + } + archiveVal := true + if len(pk.Args) > 1 { + archiveVal = resolveBool(pk.Args[1], true) + } + var update sstore.UpdatePacket + if archiveVal { + update, err = sstore.ArchiveScreen(ctx, ids.SessionId, screenId) + if err != nil { + return nil, err + } + return update, nil + } else { + log.Printf("unarchive screen %s\n", screenId) + err = sstore.UnArchiveScreen(ctx, ids.SessionId, screenId) + if err != nil { + return nil, fmt.Errorf("/screen:archive cannot un-archive screen: %v", err) + } + screen, err := sstore.GetScreenById(ctx, screenId) + if err != nil { + return nil, fmt.Errorf("/screen:archive cannot get updated screen obj: %v", err) + } + update := &sstore.ModelUpdate{ + Screens: []*sstore.ScreenType{screen}, + } + return update, nil + } +} + +func ScreenPurgeCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) // don't force R_Screen + if err != nil { + return nil, fmt.Errorf("/screen:purge cannot purge screen: %w", err) + } + screenId := ids.ScreenId + if len(pk.Args) > 0 { + ri, err := resolveSessionScreen(ctx, ids.SessionId, pk.Args[0], ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("/screen:purge cannot resolve screen arg: %v", err) + } + screenId = ri.Id + } + if screenId == "" { + return nil, fmt.Errorf("/screen:purge no active screen or screen arg passed") + } + update, err := sstore.PurgeScreen(ctx, screenId, false) + if err != nil { + return nil, err + } + return update, nil +} + +func ScreenOpenCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) + if err != nil { + return nil, fmt.Errorf("/screen:open cannot open screen: %w", err) + } + activate := resolveBool(pk.Kwargs["activate"], true) + newName := pk.Kwargs["name"] + if newName != "" { + err := validateName(newName, "screen") + if err != nil { + return nil, err + } + } + update, err := sstore.InsertScreen(ctx, ids.SessionId, newName, sstore.ScreenCreateOpts{}, activate) + if err != nil { + return nil, err + } + return update, nil +} + +func ScreenSetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + var varsUpdated []string + var setNonAnchor bool // anchor does not receive an update + updateMap := make(map[string]interface{}) + if pk.Kwargs["name"] != "" { + newName := pk.Kwargs["name"] + err = validateName(newName, "screen") + if err != nil { + return nil, err + } + updateMap[sstore.ScreenField_Name] = newName + varsUpdated = append(varsUpdated, "name") + setNonAnchor = true + } + if pk.Kwargs["sharename"] != "" { + shareName := pk.Kwargs["sharename"] + err = validateShareName(shareName) + if err != nil { + return nil, err + } + updateMap[sstore.ScreenField_ShareName] = shareName + varsUpdated = append(varsUpdated, "sharename") + setNonAnchor = true + } + if pk.Kwargs["tabcolor"] != "" { + color := pk.Kwargs["tabcolor"] + err = validateColor(color, "screen tabcolor") + if err != nil { + return nil, err + } + updateMap[sstore.ScreenField_TabColor] = color + varsUpdated = append(varsUpdated, "tabcolor") + setNonAnchor = true + } + if pk.Kwargs["pos"] != "" { + varsUpdated = append(varsUpdated, "pos") + setNonAnchor = true + } + if pk.Kwargs["focus"] != "" { + focusVal := pk.Kwargs["focus"] + if focusVal != sstore.ScreenFocusInput && focusVal != sstore.ScreenFocusCmd { + return nil, fmt.Errorf("/screen:set invalid focus argument %q, must be %s", focusVal, formatStrs([]string{sstore.ScreenFocusInput, sstore.ScreenFocusCmd}, "or", false)) + } + varsUpdated = append(varsUpdated, "focus") + updateMap[sstore.ScreenField_Focus] = focusVal + setNonAnchor = true + } + if pk.Kwargs["line"] != "" { + screen, err := sstore.GetScreenById(ctx, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("/screen:set cannot get screen: %v", err) + } + var selectedLineStr string + if screen.SelectedLine > 0 { + selectedLineStr = strconv.Itoa(int(screen.SelectedLine)) + } + ritem, err := resolveLine(ctx, screen.SessionId, screen.ScreenId, pk.Kwargs["line"], selectedLineStr) + if err != nil { + return nil, fmt.Errorf("/screen:set error resolving line: %v", err) + } + if ritem == nil { + return nil, fmt.Errorf("/screen:set could not resolve line %q", pk.Kwargs["line"]) + } + varsUpdated = append(varsUpdated, "line") + setNonAnchor = true + updateMap[sstore.ScreenField_SelectedLine] = ritem.Num + } + if pk.Kwargs["anchor"] != "" { + m := screenAnchorRe.FindStringSubmatch(pk.Kwargs["anchor"]) + if m == nil { + return nil, fmt.Errorf("/screen:set invalid anchor argument (must be [line] or [line]:[offset])") + } + anchorLine, _ := strconv.Atoi(m[1]) + varsUpdated = append(varsUpdated, "anchor") + updateMap[sstore.ScreenField_AnchorLine] = anchorLine + if m[2] != "" { + anchorOffset, _ := strconv.Atoi(m[2]) + updateMap[sstore.ScreenField_AnchorOffset] = anchorOffset + } else { + updateMap[sstore.ScreenField_AnchorOffset] = 0 + } + } + if len(varsUpdated) == 0 { + return nil, fmt.Errorf("/screen:set no updates, can set %s", formatStrs([]string{"name", "pos", "tabcolor", "focus", "anchor", "line", "sharename"}, "or", false)) + } + screen, err := sstore.UpdateScreen(ctx, ids.ScreenId, updateMap) + if err != nil { + return nil, fmt.Errorf("error updating screen: %v", err) + } + if !setNonAnchor { + return nil, nil + } + update := &sstore.ModelUpdate{ + Screens: []*sstore.ScreenType{screen}, + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("screen updated %s", formatStrs(varsUpdated, "and", false)), + TimeoutMs: 2000, + }, + } + return update, nil +} + +func ScreenCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) + if err != nil { + return nil, fmt.Errorf("/screen cannot switch to screen: %w", err) + } + firstArg := firstArg(pk) + if firstArg == "" { + return nil, fmt.Errorf("usage /screen [screen-name|screen-index|screen-id], no param specified") + } + ritem, err := resolveSessionScreen(ctx, ids.SessionId, firstArg, ids.ScreenId) + if err != nil { + return nil, err + } + update, err := sstore.SwitchScreenById(ctx, ids.SessionId, ritem.Id) + if err != nil { + return nil, err + } + return update, nil +} + +var screenAnchorRe = regexp.MustCompile("^(\\d+)(?::(-?\\d+))?$") + +func RemoteInstallCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + mshell := ids.Remote.MShell + go mshell.RunInstall() + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: ids.Remote.RemotePtr.RemoteId, + }, + }, nil +} + +func RemoteInstallCancelCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + mshell := ids.Remote.MShell + go mshell.CancelInstall() + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: ids.Remote.RemotePtr.RemoteId, + }, + }, nil +} + +func RemoteConnectCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + go ids.Remote.MShell.Launch(true) + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: ids.Remote.RemotePtr.RemoteId, + }, + }, nil +} + +func RemoteDisconnectCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + force := resolveBool(pk.Kwargs["force"], false) + go ids.Remote.MShell.Disconnect(force) + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: ids.Remote.RemotePtr.RemoteId, + }, + }, nil +} + +func makeRemoteEditUpdate_new(err error) sstore.UpdatePacket { + redit := &sstore.RemoteEditType{ + RemoteEdit: true, + } + if err != nil { + redit.ErrorStr = err.Error() + } + update := &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + RemoteEdit: redit, + }, + } + return update +} + +func makeRemoteEditErrorReturn_new(visual bool, err error) (sstore.UpdatePacket, error) { + if visual { + return makeRemoteEditUpdate_new(err), nil + } + return nil, err +} + +func makeRemoteEditUpdate_edit(ids resolvedIds, err error) sstore.UpdatePacket { + redit := &sstore.RemoteEditType{ + RemoteEdit: true, + } + redit.RemoteId = ids.Remote.RemotePtr.RemoteId + if ids.Remote.RemoteCopy.SSHOpts != nil { + redit.KeyStr = ids.Remote.RemoteCopy.SSHOpts.SSHIdentity + redit.HasPassword = (ids.Remote.RemoteCopy.SSHOpts.SSHPassword != "") + } + if err != nil { + redit.ErrorStr = err.Error() + } + update := &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + RemoteEdit: redit, + }, + } + return update +} + +func makeRemoteEditErrorReturn_edit(ids resolvedIds, visual bool, err error) (sstore.UpdatePacket, error) { + if visual { + return makeRemoteEditUpdate_edit(ids, err), nil + } + return nil, err +} + +type RemoteEditArgs struct { + CanonicalName string + SSHOpts *sstore.SSHOpts + ConnectMode string + Alias string + AutoInstall bool + SSHPassword string + SSHKeyFile string + Color string + EditMap map[string]interface{} +} + +func parseRemoteEditArgs(isNew bool, pk *scpacket.FeCommandPacketType, isLocal bool) (*RemoteEditArgs, error) { + var canonicalName string + var sshOpts *sstore.SSHOpts + var isSudo bool + + if isNew { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/remote:new must specify user@host argument (set visual=1 to edit in UI)") + } + userHost := pk.Args[0] + m := userHostRe.FindStringSubmatch(userHost) + if m == nil { + return nil, fmt.Errorf("invalid format of user@host argument") + } + sudoStr, remoteUser, remoteHost, remotePortStr := m[1], m[2], m[3], m[4] + var uhPort int + if remotePortStr != "" { + var err error + uhPort, err = strconv.Atoi(remotePortStr) + if err != nil { + return nil, fmt.Errorf("invalid port specified on user@host argument") + } + } + if sudoStr != "" { + isSudo = true + } + if pk.Kwargs["sudo"] != "" { + sudoArg := resolveBool(pk.Kwargs["sudo"], false) + if isSudo && !sudoArg { + return nil, fmt.Errorf("invalid 'sudo' argument, with sudo kw arg set to false") + } + if !isSudo && sudoArg { + isSudo = true + } + } + sshOpts = &sstore.SSHOpts{ + Local: false, + SSHHost: remoteHost, + SSHUser: remoteUser, + IsSudo: isSudo, + } + portVal, err := resolvePosInt(pk.Kwargs["port"], 0) + if err != nil { + return nil, fmt.Errorf("invalid port %q: %v", pk.Kwargs["port"], err) + } + if portVal != 0 && uhPort != 0 && portVal != uhPort { + return nil, fmt.Errorf("invalid port argument, does not match port specified in 'user@host:port' argument") + } + if portVal == 0 && uhPort != 0 { + portVal = uhPort + } + sshOpts.SSHPort = portVal + canonicalName = remoteUser + "@" + remoteHost + if isSudo { + canonicalName = "sudo@" + canonicalName + } + } else { + if pk.Kwargs["sudo"] != "" { + return nil, fmt.Errorf("cannot update 'sudo' value") + } + if pk.Kwargs["port"] != "" { + return nil, fmt.Errorf("cannot update 'port' value") + } + } + alias := pk.Kwargs["alias"] + if alias != "" { + if len(alias) > MaxRemoteAliasLen { + return nil, fmt.Errorf("alias too long, max length = %d", MaxRemoteAliasLen) + } + if !remoteAliasRe.MatchString(alias) { + return nil, fmt.Errorf("invalid alias format") + } + } + var connectMode string + if isNew { + connectMode = sstore.ConnectModeAuto + } + if pk.Kwargs["connectmode"] != "" { + connectMode = pk.Kwargs["connectmode"] + } + if connectMode != "" && !sstore.IsValidConnectMode(connectMode) { + err := fmt.Errorf("invalid connectmode %q: valid modes are %s", connectMode, formatStrs([]string{sstore.ConnectModeStartup, sstore.ConnectModeAuto, sstore.ConnectModeManual}, "or", false)) + return nil, err + } + autoInstall := resolveBool(pk.Kwargs["autoinstall"], true) + keyFile, err := resolveFile(pk.Kwargs["key"]) + if err != nil { + return nil, fmt.Errorf("invalid ssh keyfile %q: %v", pk.Kwargs["key"], err) + } + color := pk.Kwargs["color"] + if color != "" { + err := validateRemoteColor(color, "remote color") + if err != nil { + return nil, err + } + } + sshPassword := pk.Kwargs["password"] + if sshOpts != nil { + sshOpts.SSHIdentity = keyFile + sshOpts.SSHPassword = sshPassword + } + + // set up editmap + editMap := make(map[string]interface{}) + if _, found := pk.Kwargs[sstore.RemoteField_Alias]; found { + editMap[sstore.RemoteField_Alias] = alias + } + if connectMode != "" { + if isLocal { + return nil, fmt.Errorf("Cannot edit connect mode for 'local' remote") + } + editMap[sstore.RemoteField_ConnectMode] = connectMode + } + if _, found := pk.Kwargs[sstore.RemoteField_AutoInstall]; found { + editMap[sstore.RemoteField_AutoInstall] = autoInstall + } + if _, found := pk.Kwargs["key"]; found { + if isLocal { + return nil, fmt.Errorf("Cannot edit ssh key file for 'local' remote") + } + editMap[sstore.RemoteField_SSHKey] = keyFile + } + if _, found := pk.Kwargs[sstore.RemoteField_Color]; found { + editMap[sstore.RemoteField_Color] = color + } + if _, found := pk.Kwargs["password"]; found && pk.Kwargs["password"] != PasswordUnchangedSentinel { + if isLocal { + return nil, fmt.Errorf("Cannot edit ssh password for 'local' remote") + } + editMap[sstore.RemoteField_SSHPassword] = sshPassword + } + + return &RemoteEditArgs{ + SSHOpts: sshOpts, + ConnectMode: connectMode, + Alias: alias, + AutoInstall: autoInstall, + CanonicalName: canonicalName, + SSHKeyFile: keyFile, + SSHPassword: sshPassword, + Color: color, + EditMap: editMap, + }, nil +} + +func RemoteNewCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + visualEdit := resolveBool(pk.Kwargs["visual"], false) + isSubmitted := resolveBool(pk.Kwargs["submit"], false) + if visualEdit && !isSubmitted && len(pk.Args) == 0 { + return makeRemoteEditUpdate_new(nil), nil + } + editArgs, err := parseRemoteEditArgs(true, pk, false) + if err != nil { + return makeRemoteEditErrorReturn_new(visualEdit, fmt.Errorf("/remote:new %v", err)) + } + r := &sstore.RemoteType{ + RemoteId: scbase.GenPromptUUID(), + RemoteType: sstore.RemoteTypeSsh, + RemoteAlias: editArgs.Alias, + RemoteCanonicalName: editArgs.CanonicalName, + RemoteUser: editArgs.SSHOpts.SSHUser, + RemoteHost: editArgs.SSHOpts.SSHHost, + ConnectMode: editArgs.ConnectMode, + AutoInstall: editArgs.AutoInstall, + SSHOpts: editArgs.SSHOpts, + } + if editArgs.Color != "" { + r.RemoteOpts = &sstore.RemoteOptsType{Color: editArgs.Color} + } + err = remote.AddRemote(ctx, r, true) + if err != nil { + return makeRemoteEditErrorReturn_new(visualEdit, fmt.Errorf("cannot create remote %q: %v", r.RemoteCanonicalName, err)) + } + // SUCCESS + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: r.RemoteId, + }, + }, nil +} + +func RemoteSetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + visualEdit := resolveBool(pk.Kwargs["visual"], false) + isSubmitted := resolveBool(pk.Kwargs["submit"], false) + editArgs, err := parseRemoteEditArgs(false, pk, ids.Remote.MShell.IsLocal()) + if err != nil { + return makeRemoteEditErrorReturn_edit(ids, visualEdit, fmt.Errorf("/remote:new %v", err)) + } + if visualEdit && !isSubmitted && len(editArgs.EditMap) == 0 { + return makeRemoteEditUpdate_edit(ids, nil), nil + } + if !visualEdit && len(editArgs.EditMap) == 0 { + return nil, fmt.Errorf("/remote:set no updates, can set %s. (set visual=1 to edit in UI)", formatStrs(RemoteSetArgs, "or", false)) + } + err = ids.Remote.MShell.UpdateRemote(ctx, editArgs.EditMap) + if err != nil { + return makeRemoteEditErrorReturn_edit(ids, visualEdit, fmt.Errorf("/remote:new error updating remote: %v", err)) + } + if visualEdit { + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: ids.Remote.RemoteCopy.RemoteId, + }, + }, nil + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("remote %q updated", ids.Remote.DisplayName), + TimeoutMs: 2000, + }, + } + return update, nil +} + +func RemoteShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + state := ids.Remote.RState + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + PtyRemoteId: state.RemoteId, + }, + }, nil +} + +func RemoteShowAllCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + stateArr := remote.GetAllRemoteRuntimeState() + var buf bytes.Buffer + for _, rstate := range stateArr { + var name string + if rstate.RemoteAlias == "" { + name = rstate.RemoteCanonicalName + } else { + name = fmt.Sprintf("%s (%s)", rstate.RemoteCanonicalName, rstate.RemoteAlias) + } + buf.WriteString(fmt.Sprintf("%-12s %-5s %8s %s\n", rstate.Status, rstate.RemoteType, rstate.RemoteId[0:8], name)) + } + return &sstore.ModelUpdate{ + RemoteView: &sstore.RemoteViewType{ + RemoteShowAll: true, + }, + }, nil +} + +func ScreenShowAllCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) + screenArr, err := sstore.GetSessionScreens(ctx, ids.SessionId) + if err != nil { + return nil, fmt.Errorf("/screen:showall error getting screen list: %v", err) + } + var buf bytes.Buffer + for _, screen := range screenArr { + var archivedStr string + if screen.Archived { + archivedStr = " (archived)" + } + screenIdxStr := "-" + if screen.ScreenIdx != 0 { + screenIdxStr = strconv.Itoa(int(screen.ScreenIdx)) + } + outStr := fmt.Sprintf("%-30s %s %s\n", screen.Name+archivedStr, screen.ScreenId, screenIdxStr) + buf.WriteString(outStr) + } + return &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("all screens for session"), + InfoLines: splitLinesForInfo(buf.String()), + }, + }, nil +} + +func ScreenResetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + localRemote := remote.GetLocalRemote() + if localRemote == nil { + return nil, fmt.Errorf("error getting local remote (not found)") + } + rptr := sstore.RemotePtrType{RemoteId: localRemote.RemoteId} + sessionUpdate := &sstore.SessionType{SessionId: ids.SessionId} + ris, err := sstore.ScreenReset(ctx, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("error resetting screen: %v", err) + } + sessionUpdate.Remotes = append(sessionUpdate.Remotes, ris...) + err = sstore.UpdateCurRemote(ctx, ids.ScreenId, rptr) + if err != nil { + return nil, fmt.Errorf("cannot reset screen remote back to local: %w", err) + } + outputStr := "reset screen state (all remote state reset)" + cmd, err := makeStaticCmd(ctx, "screen:reset", ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update, err := addLineForCmd(ctx, "/screen:reset", false, ids, cmd, "", nil) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + update.Sessions = []*sstore.SessionType{sessionUpdate} + return update, nil +} + +func RemoteArchiveCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + err = remote.ArchiveRemote(ctx, ids.Remote.RemotePtr.RemoteId) + if err != nil { + return nil, fmt.Errorf("archiving remote: %v", err) + } + update := sstore.InfoMsgUpdate("remote [%s] archived", ids.Remote.DisplayName) + localRemote := remote.GetLocalRemote() + rptr := sstore.RemotePtrType{RemoteId: localRemote.GetRemoteId()} + err = sstore.UpdateCurRemote(ctx, ids.ScreenId, rptr) + if err != nil { + return nil, fmt.Errorf("cannot switch remote back to local: %w", err) + } + screen, err := sstore.GetScreenById(ctx, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("cannot get updated screen: %w", err) + } + update.Screens = []*sstore.ScreenType{screen} + return update, nil +} + +func RemoteCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + return nil, fmt.Errorf("/remote requires a subcommand: %s", formatStrs([]string{"show"}, "or", false)) +} + +func crShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType, ids resolvedIds) (sstore.UpdatePacket, error) { + var buf bytes.Buffer + riArr, err := sstore.GetRIsForScreen(ctx, ids.SessionId, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("cannot get remote instances: %w", err) + } + rmap := remote.GetRemoteMap() + for _, ri := range riArr { + rptr := sstore.RemotePtrType{RemoteId: ri.RemoteId, Name: ri.Name} + msh := rmap[ri.RemoteId] + if msh == nil { + continue + } + baseDisplayName := msh.GetDisplayName() + displayName := rptr.GetDisplayName(baseDisplayName) + cwdStr := "-" + if ri.FeState["cwd"] != "" { + cwdStr = ri.FeState["cwd"] + } + buf.WriteString(fmt.Sprintf("%-30s %-50s\n", displayName, cwdStr)) + } + riBaseMap := make(map[string]bool) + for _, ri := range riArr { + if ri.Name == "" { + riBaseMap[ri.RemoteId] = true + } + } + for remoteId, msh := range rmap { + if riBaseMap[remoteId] { + continue + } + feState := msh.GetDefaultFeState() + if feState == nil { + continue + } + cwdStr := "-" + if feState["cwd"] != "" { + cwdStr = feState["cwd"] + } + buf.WriteString(fmt.Sprintf("%-30s %-50s (default)\n", msh.GetDisplayName(), cwdStr)) + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func GetFullRemoteDisplayName(rptr *sstore.RemotePtrType, rstate *remote.RemoteRuntimeState) string { + if rptr == nil { + return "(invalid)" + } + if rstate.RemoteAlias != "" { + fullName := rstate.RemoteAlias + if rptr.Name != "" { + fullName = fullName + ":" + rptr.Name + } + return fmt.Sprintf("[%s] (%s)", fullName, rstate.RemoteCanonicalName) + } else { + if rptr.Name != "" { + return fmt.Sprintf("[%s:%s]", rstate.RemoteCanonicalName, rptr.Name) + } + return fmt.Sprintf("[%s]", rstate.RemoteCanonicalName) + } +} + +func writeErrorToPty(cmd *sstore.CmdType, errStr string, outputPos int64) { + errPk := openai.CreateErrorPacket(errStr) + errBytes, err := packet.MarshalPacket(errPk) + if err != nil { + log.Printf("error writing error packet to openai response: %v\n", err) + return + } + errCtx, cancelFn := context.WithTimeout(context.Background(), 5*time.Second) + defer cancelFn() + update, err := sstore.AppendToCmdPtyBlob(errCtx, cmd.ScreenId, cmd.LineId, errBytes, outputPos) + if err != nil { + log.Printf("error writing ptyupdate for openai response: %v\n", err) + return + } + sstore.MainBus.SendScreenUpdate(cmd.ScreenId, update) + return +} + +func writePacketToPty(ctx context.Context, cmd *sstore.CmdType, pk packet.PacketType, outputPos *int64) error { + outBytes, err := packet.MarshalPacket(pk) + if err != nil { + return err + } + update, err := sstore.AppendToCmdPtyBlob(ctx, cmd.ScreenId, cmd.LineId, outBytes, *outputPos) + if err != nil { + return err + } + *outputPos += int64(len(outBytes)) + sstore.MainBus.SendScreenUpdate(cmd.ScreenId, update) + return nil +} + +func doOpenAICompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) { + var outputPos int64 + var hadError bool + startTime := time.Now() + ctx, cancelFn := context.WithTimeout(context.Background(), 30*time.Second) + defer cancelFn() + defer func() { + r := recover() + if r != nil { + panicMsg := fmt.Sprintf("panic: %v", r) + log.Printf("panic in doOpenAICompletion: %s\n", panicMsg) + writeErrorToPty(cmd, panicMsg, outputPos) + hadError = true + } + duration := time.Since(startTime) + cmdStatus := sstore.CmdStatusDone + var exitCode int + if hadError { + cmdStatus = sstore.CmdStatusError + exitCode = 1 + } + ck := base.MakeCommandKey(cmd.ScreenId, cmd.LineId) + donePk := packet.MakeCmdDonePacket(ck) + donePk.Ts = time.Now().UnixMilli() + donePk.ExitCode = exitCode + donePk.DurationMs = duration.Milliseconds() + update, err := sstore.UpdateCmdDoneInfo(context.Background(), ck, donePk, cmdStatus) + if err != nil { + // nothing to do + log.Printf("error updating cmddoneinfo (in openai): %v\n", err) + return + } + sstore.MainBus.SendScreenUpdate(cmd.ScreenId, update) + }() + respPks, err := openai.RunCompletion(ctx, opts, prompt) + if err != nil { + writeErrorToPty(cmd, fmt.Sprintf("error calling OpenAI API: %v", err), outputPos) + return + } + for _, pk := range respPks { + err = writePacketToPty(ctx, cmd, pk, &outputPos) + if err != nil { + writeErrorToPty(cmd, fmt.Sprintf("error writing response to ptybuffer: %v", err), outputPos) + return + } + } + return +} + +func doOpenAIStreamCompletion(cmd *sstore.CmdType, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) { + var outputPos int64 + var hadError bool + startTime := time.Now() + ctx, cancelFn := context.WithTimeout(context.Background(), 30*time.Second) + defer cancelFn() + defer func() { + r := recover() + if r != nil { + panicMsg := fmt.Sprintf("panic: %v", r) + log.Printf("panic in doOpenAICompletion: %s\n", panicMsg) + writeErrorToPty(cmd, panicMsg, outputPos) + hadError = true + } + duration := time.Since(startTime) + cmdStatus := sstore.CmdStatusDone + var exitCode int + if hadError { + cmdStatus = sstore.CmdStatusError + exitCode = 1 + } + ck := base.MakeCommandKey(cmd.ScreenId, cmd.LineId) + donePk := packet.MakeCmdDonePacket(ck) + donePk.Ts = time.Now().UnixMilli() + donePk.ExitCode = exitCode + donePk.DurationMs = duration.Milliseconds() + update, err := sstore.UpdateCmdDoneInfo(context.Background(), ck, donePk, cmdStatus) + if err != nil { + // nothing to do + log.Printf("error updating cmddoneinfo (in openai): %v\n", err) + return + } + sstore.MainBus.SendScreenUpdate(cmd.ScreenId, update) + }() + ch, err := openai.RunCompletionStream(ctx, opts, prompt) + if err != nil { + writeErrorToPty(cmd, fmt.Sprintf("error calling OpenAI API: %v", err), outputPos) + return + } + for pk := range ch { + err = writePacketToPty(ctx, cmd, pk, &outputPos) + if err != nil { + writeErrorToPty(cmd, fmt.Sprintf("error writing response to ptybuffer: %v", err), outputPos) + return + } + } + return +} + +func OpenAICommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, fmt.Errorf("/%s error: %w", GetCmdStr(pk), err) + } + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + if clientData.OpenAIOpts == nil || clientData.OpenAIOpts.APIToken == "" { + return nil, fmt.Errorf("no openai API token found, configure in client settings") + } + opts := clientData.OpenAIOpts + if opts.Model == "" { + opts.Model = openai.DefaultModel + } + if opts.MaxTokens == 0 { + opts.MaxTokens = openai.DefaultMaxTokens + } + promptStr := firstArg(pk) + if promptStr == "" { + return nil, fmt.Errorf("openai error, prompt string is blank") + } + ptermVal := defaultStr(pk.Kwargs["pterm"], DefaultPTERM) + pkTermOpts, err := GetUITermOpts(pk.UIContext.WinSize, ptermVal) + if err != nil { + return nil, fmt.Errorf("openai error, invalid 'pterm' value %q: %v", ptermVal, err) + } + termOpts := convertTermOpts(pkTermOpts) + cmd, err := makeDynCmd(ctx, GetCmdStr(pk), ids, pk.GetRawStr(), *termOpts) + if err != nil { + return nil, fmt.Errorf("openai error, cannot make dyn cmd") + } + line, err := sstore.AddOpenAILine(ctx, ids.ScreenId, DefaultUserId, cmd) + if err != nil { + return nil, fmt.Errorf("cannot add new line: %v", err) + } + prompt := []sstore.OpenAIPromptMessageType{{Role: sstore.OpenAIRoleUser, Content: promptStr}} + if resolveBool(pk.Kwargs["stream"], true) { + go doOpenAIStreamCompletion(cmd, opts, prompt) + } else { + go doOpenAICompletion(cmd, opts, prompt) + } + updateHistoryContext(ctx, line, cmd) + updateMap := make(map[string]interface{}) + updateMap[sstore.ScreenField_SelectedLine] = line.LineNum + updateMap[sstore.ScreenField_Focus] = sstore.ScreenFocusInput + screen, err := sstore.UpdateScreen(ctx, ids.ScreenId, updateMap) + if err != nil { + // ignore error again (nothing to do) + log.Printf("openai error updating screen selected line: %v\n", err) + } + update := &sstore.ModelUpdate{Line: line, Cmd: cmd, Screens: []*sstore.ScreenType{screen}} + return update, nil +} + +func CrCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, fmt.Errorf("/%s error: %w", GetCmdStr(pk), err) + } + newRemote := firstArg(pk) + if newRemote == "" { + return crShowCommand(ctx, pk, ids) + } + _, rptr, rstate, err := resolveRemote(ctx, newRemote, ids.SessionId, ids.ScreenId) + if err != nil { + return nil, err + } + if rptr == nil { + return nil, fmt.Errorf("/%s error: remote %q not found", GetCmdStr(pk), newRemote) + } + if rstate.Archived { + return nil, fmt.Errorf("/%s error: remote %q cannot switch to archived remote", GetCmdStr(pk), newRemote) + } + err = sstore.UpdateCurRemote(ctx, ids.ScreenId, *rptr) + if err != nil { + return nil, fmt.Errorf("/%s error: cannot update curremote: %w", GetCmdStr(pk), err) + } + outputStr := fmt.Sprintf("connected to %s", GetFullRemoteDisplayName(rptr, rstate)) + cmd, err := makeStaticCmd(ctx, GetCmdStr(pk), ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update, err := addLineForCmd(ctx, "/"+GetCmdStr(pk), false, ids, cmd, "", nil) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + return update, nil +} + +func makeDynCmd(ctx context.Context, metaCmd string, ids resolvedIds, cmdStr string, termOpts sstore.TermOpts) (*sstore.CmdType, error) { + cmd := &sstore.CmdType{ + ScreenId: ids.ScreenId, + LineId: scbase.GenPromptUUID(), + CmdStr: cmdStr, + RawCmdStr: cmdStr, + Remote: ids.Remote.RemotePtr, + TermOpts: termOpts, + Status: sstore.CmdStatusRunning, + RunOut: nil, + } + if ids.Remote.StatePtr != nil { + cmd.StatePtr = *ids.Remote.StatePtr + } + if ids.Remote.FeState != nil { + cmd.FeState = ids.Remote.FeState + } + err := sstore.CreateCmdPtyFile(ctx, cmd.ScreenId, cmd.LineId, cmd.TermOpts.MaxPtySize) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, fmt.Errorf("cannot create local ptyout file for %s command: %w", metaCmd, err) + } + return cmd, nil +} + +func makeStaticCmd(ctx context.Context, metaCmd string, ids resolvedIds, cmdStr string, cmdOutput []byte) (*sstore.CmdType, error) { + cmd := &sstore.CmdType{ + ScreenId: ids.ScreenId, + LineId: scbase.GenPromptUUID(), + CmdStr: cmdStr, + RawCmdStr: cmdStr, + Remote: ids.Remote.RemotePtr, + TermOpts: sstore.TermOpts{Rows: shexec.DefaultTermRows, Cols: shexec.DefaultTermCols, FlexRows: true, MaxPtySize: remote.DefaultMaxPtySize}, + Status: sstore.CmdStatusDone, + RunOut: nil, + } + if ids.Remote.StatePtr != nil { + cmd.StatePtr = *ids.Remote.StatePtr + } + if ids.Remote.FeState != nil { + cmd.FeState = ids.Remote.FeState + } + err := sstore.CreateCmdPtyFile(ctx, cmd.ScreenId, cmd.LineId, cmd.TermOpts.MaxPtySize) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, fmt.Errorf("cannot create local ptyout file for %s command: %w", metaCmd, err) + } + // can ignore ptyupdate + _, err = sstore.AppendToCmdPtyBlob(ctx, ids.ScreenId, cmd.LineId, cmdOutput, 0) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, fmt.Errorf("cannot append to local ptyout file for %s command: %v", metaCmd, err) + } + return cmd, nil +} + +func addLineForCmd(ctx context.Context, metaCmd string, shouldFocus bool, ids resolvedIds, cmd *sstore.CmdType, renderer string, lineState map[string]any) (*sstore.ModelUpdate, error) { + rtnLine, err := sstore.AddCmdLine(ctx, ids.ScreenId, DefaultUserId, cmd, renderer, lineState) + if err != nil { + return nil, err + } + screen, err := sstore.GetScreenById(ctx, ids.ScreenId) + if err != nil { + // ignore error here, because the command has already run (nothing to do) + log.Printf("%s error getting screen: %v\n", metaCmd, err) + } + if screen != nil { + updateMap := make(map[string]interface{}) + updateMap[sstore.ScreenField_SelectedLine] = rtnLine.LineNum + if shouldFocus { + updateMap[sstore.ScreenField_Focus] = sstore.ScreenFocusCmd + } + screen, err = sstore.UpdateScreen(ctx, ids.ScreenId, updateMap) + if err != nil { + // ignore error again (nothing to do) + log.Printf("%s error updating screen selected line: %v\n", metaCmd, err) + } + } + update := &sstore.ModelUpdate{ + Line: rtnLine, + Cmd: cmd, + Screens: []*sstore.ScreenType{screen}, + } + updateHistoryContext(ctx, rtnLine, cmd) + return update, nil +} + +func updateHistoryContext(ctx context.Context, line *sstore.LineType, cmd *sstore.CmdType) { + ctxVal := ctx.Value(historyContextKey) + if ctxVal == nil { + return + } + hctx := ctxVal.(*historyContextType) + if line != nil { + hctx.LineId = line.LineId + hctx.LineNum = line.LineNum + } + if cmd != nil { + hctx.RemotePtr = &cmd.Remote + } +} + +func makeInfoFromComps(compType string, comps []string, hasMore bool) sstore.UpdatePacket { + sort.Slice(comps, func(i int, j int) bool { + c1 := comps[i] + c2 := comps[j] + c1mc := strings.HasPrefix(c1, "^") + c2mc := strings.HasPrefix(c2, "^") + if c1mc && !c2mc { + return true + } + if !c1mc && c2mc { + return false + } + return c1 < c2 + }) + if len(comps) == 0 { + comps = []string{"(no completions)"} + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("%s completions", compType), + InfoComps: comps, + InfoCompsMore: hasMore, + }, + } + return update +} + +func simpleCompCommandMeta(ctx context.Context, prefix string, compCtx comp.CompContext, args []interface{}) (*comp.CompReturn, error) { + if strings.HasPrefix(prefix, "/") { + compsCmd, _ := comp.DoSimpleComp(ctx, comp.CGTypeCommand, prefix, compCtx, nil) + compsMeta, _ := simpleCompMeta(ctx, prefix, compCtx, nil) + return comp.CombineCompReturn(comp.CGTypeCommandMeta, compsCmd, compsMeta), nil + } else { + return comp.DoSimpleComp(ctx, comp.CGTypeCommand, prefix, compCtx, nil) + } +} + +func simpleCompMeta(ctx context.Context, prefix string, compCtx comp.CompContext, args []interface{}) (*comp.CompReturn, error) { + rtn := comp.CompReturn{} + validCommands := getValidCommands() + for _, cmd := range validCommands { + if strings.HasPrefix(cmd, "/_") && !strings.HasPrefix(prefix, "/_") { + continue + } + if strings.HasPrefix(cmd, prefix) { + rtn.Entries = append(rtn.Entries, comp.CompEntry{Word: cmd, IsMetaCmd: true}) + } + } + return &rtn, nil +} + +func doMetaCompGen(ctx context.Context, pk *scpacket.FeCommandPacketType, prefix string, forDisplay bool) ([]string, bool, error) { + ids, err := resolveUiIds(ctx, pk, 0) // best effort + var comps []string + var hasMore bool + if ids.Remote != nil && ids.Remote.RState.IsConnected() { + comps, hasMore, err = doCompGen(ctx, pk, prefix, "file", forDisplay) + if err != nil { + return nil, false, err + } + } + validCommands := getValidCommands() + for _, cmd := range validCommands { + if strings.HasPrefix(cmd, prefix) { + if forDisplay { + comps = append(comps, "^"+cmd) + } else { + comps = append(comps, cmd) + } + } + } + return comps, hasMore, nil +} + +func doCompGen(ctx context.Context, pk *scpacket.FeCommandPacketType, prefix string, compType string, forDisplay bool) ([]string, bool, error) { + if compType == "metacommand" { + return doMetaCompGen(ctx, pk, prefix, forDisplay) + } + if !packet.IsValidCompGenType(compType) { + return nil, false, fmt.Errorf("/_compgen invalid type '%s'", compType) + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, false, fmt.Errorf("/_compgen error: %w", err) + } + cgPacket := packet.MakeCompGenPacket() + cgPacket.ReqId = uuid.New().String() + cgPacket.CompType = compType + cgPacket.Prefix = prefix + cgPacket.Cwd = ids.Remote.FeState["cwd"] + resp, err := ids.Remote.MShell.PacketRpc(ctx, cgPacket) + if err != nil { + return nil, false, err + } + if err = resp.Err(); err != nil { + return nil, false, err + } + comps := utilfn.GetStrArr(resp.Data, "comps") + hasMore := utilfn.GetBool(resp.Data, "hasmore") + return comps, hasMore, nil +} + +func CompGenCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, 0) // best-effort + if err != nil { + return nil, fmt.Errorf("/_compgen error: %w", err) + } + cmdLine := firstArg(pk) + pos := len(cmdLine) + if pk.Kwargs["comppos"] != "" { + posArg, err := strconv.Atoi(pk.Kwargs["comppos"]) + if err != nil { + return nil, fmt.Errorf("/_compgen invalid comppos '%s': %w", pk.Kwargs["comppos"], err) + } + pos = posArg + } + if pos < 0 { + pos = 0 + } + if pos > len(cmdLine) { + pos = len(cmdLine) + } + showComps := resolveBool(pk.Kwargs["compshow"], false) + cmdSP := utilfn.StrWithPos{Str: cmdLine, Pos: pos} + compCtx := comp.CompContext{} + if ids.Remote != nil { + rptr := ids.Remote.RemotePtr + compCtx.RemotePtr = &rptr + if ids.Remote.FeState != nil { + compCtx.Cwd = ids.Remote.FeState["cwd"] + } + } + compCtx.ForDisplay = showComps + crtn, newSP, err := comp.DoCompGen(ctx, cmdSP, compCtx) + if err != nil { + return nil, err + } + if crtn == nil { + return nil, nil + } + if showComps { + compStrs := crtn.GetCompDisplayStrs() + return makeInfoFromComps(crtn.CompType, compStrs, crtn.HasMore), nil + } + if newSP == nil || cmdSP == *newSP { + return nil, nil + } + update := &sstore.ModelUpdate{ + CmdLine: &sstore.CmdLineType{CmdLine: newSP.Str, CursorPos: newSP.Pos}, + } + return update, nil +} + +func CommentCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, fmt.Errorf("/comment error: %w", err) + } + text := firstArg(pk) + if strings.TrimSpace(text) == "" { + return nil, fmt.Errorf("cannot post empty comment") + } + rtnLine, err := sstore.AddCommentLine(ctx, ids.ScreenId, DefaultUserId, text) + if err != nil { + return nil, err + } + updateHistoryContext(ctx, rtnLine, nil) + updateMap := make(map[string]interface{}) + updateMap[sstore.ScreenField_SelectedLine] = rtnLine.LineNum + updateMap[sstore.ScreenField_Focus] = sstore.ScreenFocusInput + screen, err := sstore.UpdateScreen(ctx, ids.ScreenId, updateMap) + if err != nil { + // ignore error again (nothing to do) + log.Printf("/comment error updating screen selected line: %v\n", err) + } + update := &sstore.ModelUpdate{Line: rtnLine, Screens: []*sstore.ScreenType{screen}} + return update, nil +} + +func maybeQuote(s string, quote bool) string { + if quote { + return fmt.Sprintf("%q", s) + } + return s +} + +func mapToStrs(m map[string]bool) []string { + var rtn []string + for key, val := range m { + if val { + rtn = append(rtn, key) + } + } + return rtn +} + +func formatStrs(strs []string, conj string, quote bool) string { + if len(strs) == 0 { + return "(none)" + } + if len(strs) == 1 { + return maybeQuote(strs[0], quote) + } + if len(strs) == 2 { + return fmt.Sprintf("%s %s %s", maybeQuote(strs[0], quote), conj, maybeQuote(strs[1], quote)) + } + var buf bytes.Buffer + for idx := 0; idx < len(strs)-1; idx++ { + buf.WriteString(maybeQuote(strs[idx], quote)) + buf.WriteString(", ") + } + buf.WriteString(conj) + buf.WriteString(" ") + buf.WriteString(maybeQuote(strs[len(strs)-1], quote)) + return buf.String() +} + +func validateName(name string, typeStr string) error { + if len(name) > MaxNameLen { + return fmt.Errorf("%s name too long, max length is %d", typeStr, MaxNameLen) + } + if !genericNameRe.MatchString(name) { + return fmt.Errorf("invalid %s name", typeStr) + } + return nil +} + +func validateShareName(name string) error { + if len(name) > MaxShareNameLen { + return fmt.Errorf("share name too long, max length is %d", MaxShareNameLen) + } + for _, ch := range name { + if !unicode.IsPrint(ch) { + return fmt.Errorf("invalid character %q in share name", string(ch)) + } + } + return nil +} + +func validateRenderer(renderer string) error { + if renderer == "" { + return nil + } + if len(renderer) > MaxRendererLen { + return fmt.Errorf("renderer name too long, max length is %d", MaxRendererLen) + } + if !rendererRe.MatchString(renderer) { + return fmt.Errorf("invalid renderer format") + } + return nil +} + +func validateColor(color string, typeStr string) error { + for _, c := range ColorNames { + if color == c { + return nil + } + } + return fmt.Errorf("invalid %s, valid colors are: %s", typeStr, formatStrs(ColorNames, "or", false)) +} + +func validateRemoteColor(color string, typeStr string) error { + for _, c := range RemoteColorNames { + if color == c { + return nil + } + } + return fmt.Errorf("invalid %s, valid colors are: %s", typeStr, formatStrs(RemoteColorNames, "or", false)) +} + +func SessionOpenSharedCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + activity := sstore.ActivityUpdate{ClickShared: 1} + err := sstore.UpdateCurrentActivity(ctx, activity) + if err != nil { + log.Printf("error updating click-shared: %v\n", err) + } + return nil, fmt.Errorf("shared sessions are not available in this version of prompt (stay tuned)") +} + +func SessionOpenCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + activate := resolveBool(pk.Kwargs["activate"], true) + newName := pk.Kwargs["name"] + if newName != "" { + err := validateName(newName, "session") + if err != nil { + return nil, err + } + } + update, err := sstore.InsertSessionWithName(ctx, newName, activate) + if err != nil { + return nil, err + } + return update, nil +} + +func makeExternLink(urlStr string) string { + return fmt.Sprintf(`https://extern?%s`, url.QueryEscape(urlStr)) +} + +func ScreenWebShareCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + return nil, fmt.Errorf("websharing is no longer available") +} + +func SessionDeleteCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, 0) // don't force R_Session + if err != nil { + return nil, err + } + sessionId := "" + if len(pk.Args) >= 1 { + ritem, err := resolveSession(ctx, pk.Args[0], ids.SessionId) + if err != nil { + return nil, fmt.Errorf("/session:purge error resolving session %q: %w", pk.Args[0], err) + } + if ritem == nil { + return nil, fmt.Errorf("/session:purge session %q not found", pk.Args[0]) + } + sessionId = ritem.Id + } else { + sessionId = ids.SessionId + } + if sessionId == "" { + return nil, fmt.Errorf("/session:purge no sessionid found") + } + update, err := sstore.PurgeSession(ctx, sessionId) + if err != nil { + return nil, fmt.Errorf("cannot delete session: %v", err) + } + return update, nil +} + +func SessionArchiveCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, 0) // don't force R_Session + if err != nil { + return nil, err + } + sessionId := "" + if len(pk.Args) >= 1 { + ritem, err := resolveSession(ctx, pk.Args[0], ids.SessionId) + if err != nil { + return nil, fmt.Errorf("/session:archive error resolving session %q: %w", pk.Args[0], err) + } + if ritem == nil { + return nil, fmt.Errorf("/session:archive session %q not found", pk.Args[0]) + } + sessionId = ritem.Id + } else { + sessionId = ids.SessionId + } + if sessionId == "" { + return nil, fmt.Errorf("/session:archive no sessionid found") + } + archiveVal := true + if len(pk.Args) >= 2 { + archiveVal = resolveBool(pk.Args[1], true) + } + if archiveVal { + update, err := sstore.ArchiveSession(ctx, sessionId) + if err != nil { + return nil, fmt.Errorf("cannot archive session: %v", err) + } + update.Info = &sstore.InfoMsgType{ + InfoMsg: "session archived", + } + return update, nil + } else { + activate := resolveBool(pk.Kwargs["activate"], false) + update, err := sstore.UnArchiveSession(ctx, sessionId, activate) + if err != nil { + return nil, fmt.Errorf("cannot un-archive session: %v", err) + } + update.Info = &sstore.InfoMsgType{ + InfoMsg: "session un-archived", + } + return update, nil + } +} + +func SessionShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) + if err != nil { + return nil, err + } + session, err := sstore.GetSessionById(ctx, ids.SessionId) + if err != nil { + return nil, fmt.Errorf("cannot get session: %w", err) + } + if session == nil { + return nil, fmt.Errorf("session not found") + } + var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "sessionid", session.SessionId)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "name", session.Name)) + if session.SessionIdx != 0 { + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "index", session.SessionIdx)) + } + if session.Archived { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "archived", "true")) + ts := time.UnixMilli(session.ArchivedTs) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "archivedts", ts.Format(TsFormatStr))) + } + stats, err := sstore.GetSessionStats(ctx, ids.SessionId) + if err != nil { + return nil, fmt.Errorf("error getting session stats: %w", err) + } + var screenArchiveStr string + if stats.NumArchivedScreens > 0 { + screenArchiveStr = fmt.Sprintf(" (%d archived)", stats.NumArchivedScreens) + } + buf.WriteString(fmt.Sprintf(" %-15s %d%s\n", "screens", stats.NumScreens, screenArchiveStr)) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "lines", stats.NumLines)) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "cmds", stats.NumCmds)) + buf.WriteString(fmt.Sprintf(" %-15s %0.2fM\n", "disksize", float64(stats.DiskStats.TotalSize)/1000000)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "disk-location", stats.DiskStats.Location)) + return &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: "session info", + InfoLines: splitLinesForInfo(buf.String()), + }, + }, nil +} + +func SessionShowAllCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + sessions, err := sstore.GetBareSessions(ctx) + if err != nil { + return nil, fmt.Errorf("error retrieving sessions: %v", err) + } + var buf bytes.Buffer + for _, session := range sessions { + var archivedStr string + if session.Archived { + archivedStr = " (archived)" + } + sessionIdxStr := "-" + if session.SessionIdx != 0 { + sessionIdxStr = strconv.Itoa(int(session.SessionIdx)) + } + outStr := fmt.Sprintf("%-30s %s %s\n", session.Name+archivedStr, session.SessionId, sessionIdxStr) + buf.WriteString(outStr) + } + return &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: "all sessions", + InfoLines: splitLinesForInfo(buf.String()), + }, + }, nil +} + +func SessionSetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session) + if err != nil { + return nil, err + } + var varsUpdated []string + if pk.Kwargs["name"] != "" { + newName := pk.Kwargs["name"] + err = validateName(newName, "session") + if err != nil { + return nil, err + } + err = sstore.SetSessionName(ctx, ids.SessionId, newName) + if err != nil { + return nil, fmt.Errorf("setting session name: %v", err) + } + varsUpdated = append(varsUpdated, "name") + } + if pk.Kwargs["pos"] != "" { + + } + if len(varsUpdated) == 0 { + return nil, fmt.Errorf("/session:set no updates, can set %s", formatStrs([]string{"name", "pos"}, "or", false)) + } + bareSession, err := sstore.GetBareSessionById(ctx, ids.SessionId) + update := &sstore.ModelUpdate{ + Sessions: []*sstore.SessionType{bareSession}, + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("session updated %s", formatStrs(varsUpdated, "and", false)), + TimeoutMs: 2000, + }, + } + return update, nil +} + +func SessionCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, 0) + if err != nil { + return nil, err + } + firstArg := firstArg(pk) + if firstArg == "" { + return nil, fmt.Errorf("usage /session [name|id|pos], no param specified") + } + ritem, err := resolveSession(ctx, firstArg, ids.SessionId) + if err != nil { + return nil, err + } + err = sstore.SetActiveSessionId(ctx, ritem.Id) + if err != nil { + return nil, err + } + update := &sstore.ModelUpdate{ + ActiveSessionId: ritem.Id, + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("switched to session %q", ritem.Name), + TimeoutMs: 2000, + }, + } + return update, nil +} + +func RemoteResetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + initPk, err := ids.Remote.MShell.ReInit(ctx) + if err != nil { + return nil, err + } + if initPk == nil || initPk.State == nil { + return nil, fmt.Errorf("invalid initpk received from remote (no remote state)") + } + feState := sstore.FeStateFromShellState(initPk.State) + remoteInst, err := sstore.UpdateRemoteState(ctx, ids.SessionId, ids.ScreenId, ids.Remote.RemotePtr, feState, initPk.State, nil) + if err != nil { + return nil, err + } + outputStr := "reset remote state" + cmd, err := makeStaticCmd(ctx, "reset", ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update, err := addLineForCmd(ctx, "/reset", false, ids, cmd, "", nil) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + update.Sessions = sstore.MakeSessionsUpdateForRemote(ids.SessionId, remoteInst) + return update, nil +} + +func ClearCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if resolveBool(pk.Kwargs["purge"], false) { + update, err := sstore.PurgeScreenLines(ctx, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("clearing screen: %v", err) + } + update.Info = &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("screen cleared (all lines purged)"), + TimeoutMs: 2000, + } + return update, nil + } else { + update, err := sstore.ArchiveScreenLines(ctx, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("clearing screen: %v", err) + } + update.Info = &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("screen cleared"), + TimeoutMs: 2000, + } + return update, nil + } + +} + +func HistoryPurgeCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/history:purge requires at least one argument (history id)") + } + var historyIds []string + for _, historyArg := range pk.Args { + _, err := uuid.Parse(historyArg) + if err != nil { + return nil, fmt.Errorf("invalid historyid (must be uuid)") + } + historyIds = append(historyIds, historyArg) + } + historyItemsRemoved, err := sstore.PurgeHistoryByIds(ctx, historyIds) + if err != nil { + return nil, fmt.Errorf("/history:purge error purging items: %v", err) + } + update := &sstore.ModelUpdate{} + for _, historyItem := range historyItemsRemoved { + if historyItem.LineId == "" { + continue + } + lineObj := &sstore.LineType{ + ScreenId: historyItem.ScreenId, + LineId: historyItem.LineId, + Remove: true, + } + update.Lines = append(update.Lines, lineObj) + } + return update, nil +} + +const HistoryViewPageSize = 50 + +var cmdFilterLs = regexp.MustCompile(`^ls(\s|$)`) +var cmdFilterCd = regexp.MustCompile(`^cd(\s|$)`) + +func historyCmdFilter(hitem *sstore.HistoryItemType) bool { + cmdStr := hitem.CmdStr + if cmdStr == "" || strings.Index(cmdStr, ";") != -1 || strings.Index(cmdStr, "\n") != -1 { + return true + } + if cmdFilterLs.MatchString(cmdStr) { + return false + } + if cmdFilterCd.MatchString(cmdStr) { + return false + } + return true +} + +func HistoryViewAllCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + _, err := resolveUiIds(ctx, pk, 0) + if err != nil { + return nil, err + } + offset, err := resolveNonNegInt(pk.Kwargs["offset"], 0) + if err != nil { + return nil, err + } + rawOffset, err := resolveNonNegInt(pk.Kwargs["rawoffset"], 0) + if err != nil { + return nil, err + } + opts := sstore.HistoryQueryOpts{MaxItems: HistoryViewPageSize, Offset: offset, RawOffset: rawOffset} + if pk.Kwargs["text"] != "" { + opts.SearchText = pk.Kwargs["text"] + } + if pk.Kwargs["searchsession"] != "" { + sessionId, err := resolveSessionArg(pk.Kwargs["searchsession"]) + if err != nil { + return nil, fmt.Errorf("invalid searchsession: %v", err) + } + opts.SessionId = sessionId + } + if pk.Kwargs["searchremote"] != "" { + rptr, err := resolveRemoteArg(pk.Kwargs["searchremote"]) + if err != nil { + return nil, fmt.Errorf("invalid searchremote: %v", err) + } + if rptr != nil { + opts.RemoteId = rptr.RemoteId + } + } + if pk.Kwargs["fromts"] != "" { + fromTs, err := resolvePosInt(pk.Kwargs["fromts"], 0) + if err != nil { + return nil, fmt.Errorf("invalid fromts (must be unixtime (milliseconds): %v", err) + } + if fromTs > 0 { + opts.FromTs = int64(fromTs) + } + } + if pk.Kwargs["meta"] != "" { + opts.NoMeta = !resolveBool(pk.Kwargs["meta"], true) + } + if resolveBool(pk.Kwargs["filter"], false) { + opts.FilterFn = historyCmdFilter + } + if err != nil { + return nil, fmt.Errorf("invalid meta arg (must be boolean): %v", err) + } + hresult, err := sstore.GetHistoryItems(ctx, opts) + if err != nil { + return nil, err + } + hvdata := &sstore.HistoryViewData{ + Items: hresult.Items, + Offset: hresult.Offset, + RawOffset: hresult.RawOffset, + NextRawOffset: hresult.NextRawOffset, + HasMore: hresult.HasMore, + } + lines, cmds, err := sstore.GetLineCmdsFromHistoryItems(ctx, hvdata.Items) + if err != nil { + return nil, err + } + hvdata.Lines = lines + hvdata.Cmds = cmds + update := &sstore.ModelUpdate{ + HistoryViewData: hvdata, + MainView: sstore.MainViewHistory, + } + return update, nil +} + +const DefaultMaxHistoryItems = 10000 + +func HistoryCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_Remote) + if err != nil { + return nil, err + } + maxItems, err := resolvePosInt(pk.Kwargs["maxitems"], DefaultMaxHistoryItems) + if err != nil { + return nil, fmt.Errorf("invalid maxitems value '%s' (must be a number): %v", pk.Kwargs["maxitems"], err) + } + if maxItems < 0 { + return nil, fmt.Errorf("invalid maxitems value '%d' (cannot be negative)", maxItems) + } + if maxItems == 0 { + maxItems = DefaultMaxHistoryItems + } + htype := HistoryTypeScreen + hSessionId := ids.SessionId + hScreenId := ids.ScreenId + if pk.Kwargs["type"] != "" { + htype = pk.Kwargs["type"] + if htype != HistoryTypeScreen && htype != HistoryTypeSession && htype != HistoryTypeGlobal { + return nil, fmt.Errorf("invalid history type '%s', valid types: %s", htype, formatStrs([]string{HistoryTypeScreen, HistoryTypeSession, HistoryTypeGlobal}, "or", false)) + } + } + if htype == HistoryTypeGlobal { + hSessionId = "" + hScreenId = "" + } else if htype == HistoryTypeSession { + hScreenId = "" + } + hopts := sstore.HistoryQueryOpts{MaxItems: maxItems, SessionId: hSessionId, ScreenId: hScreenId} + hresult, err := sstore.GetHistoryItems(ctx, hopts) + if err != nil { + return nil, err + } + show := !resolveBool(pk.Kwargs["noshow"], false) + if show { + err = sstore.UpdateCurrentActivity(ctx, sstore.ActivityUpdate{HistoryView: 1}) + if err != nil { + log.Printf("error updating current activity (history): %v\n", err) + } + } + update := &sstore.ModelUpdate{} + update.History = &sstore.HistoryInfoType{ + HistoryType: htype, + SessionId: ids.SessionId, + ScreenId: ids.ScreenId, + Items: hresult.Items, + Show: show, + } + return update, nil +} + +func splitLinesForInfo(str string) []string { + rtn := strings.Split(str, "\n") + if rtn[len(rtn)-1] == "" { + return rtn[:len(rtn)-1] + } + return rtn +} + +func resizeRunningCommand(ctx context.Context, cmd *sstore.CmdType, newCols int) error { + siPk := packet.MakeSpecialInputPacket() + siPk.CK = base.MakeCommandKey(cmd.ScreenId, cmd.LineId) + siPk.WinSize = &packet.WinSize{Rows: int(cmd.TermOpts.Rows), Cols: newCols} + msh := remote.GetRemoteById(cmd.Remote.RemoteId) + if msh == nil { + return fmt.Errorf("cannot resize, cmd remote not found") + } + err := msh.SendSpecialInput(siPk) + if err != nil { + return err + } + newTermOpts := cmd.TermOpts + newTermOpts.Cols = int64(newCols) + err = sstore.UpdateCmdTermOpts(ctx, cmd.ScreenId, cmd.LineId, newTermOpts) + if err != nil { + return err + } + return nil +} + +func ScreenResizeCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + colsStr := pk.Kwargs["cols"] + if colsStr == "" { + return nil, fmt.Errorf("/screen:resize requires a numeric 'cols' argument") + } + cols, err := strconv.Atoi(colsStr) + if err != nil { + return nil, fmt.Errorf("/screen:resize requires a numeric 'cols' argument: %v", err) + } + if cols <= 0 { + return nil, fmt.Errorf("/screen:resize invalid zero/negative 'cols' argument") + } + cols = base.BoundInt(cols, shexec.MinTermCols, shexec.MaxTermCols) + runningCmds, err := sstore.GetRunningScreenCmds(ctx, ids.ScreenId) + if err != nil { + return nil, fmt.Errorf("/screen:resize cannot get running commands: %v", err) + } + if len(runningCmds) == 0 { + return nil, nil + } + for _, cmd := range runningCmds { + if int(cmd.TermOpts.Cols) != cols { + resizeRunningCommand(ctx, cmd, cols) + } + } + return nil, nil +} + +func LineCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + return nil, fmt.Errorf("/line requires a subcommand: %s", formatStrs([]string{"show", "star", "hide", "purge", "setheight", "set"}, "or", false)) +} + +func LineSetHeightCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) != 2 { + return nil, fmt.Errorf("/line:setheight requires 2 arguments (linearg and height)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + heightVal, err := resolveNonNegInt(pk.Args[1], 0) + if err != nil { + return nil, fmt.Errorf("/line:setheight invalid height val: %v", err) + } + if heightVal > 10000 { + return nil, fmt.Errorf("/line:setheight invalid height val (too large): %d", heightVal) + } + err = sstore.UpdateLineHeight(ctx, ids.ScreenId, lineId, heightVal) + if err != nil { + return nil, fmt.Errorf("/line:setheight error updating height: %v", err) + } + // we don't need to pass the updated line height (it is "write only") + return nil, nil +} + +func LineSetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) != 1 { + return nil, fmt.Errorf("/line:set requires 1 argument (linearg)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + var varsUpdated []string + if renderer, found := pk.Kwargs[KwArgRenderer]; found { + if err = validateRenderer(renderer); err != nil { + return nil, fmt.Errorf("invalid renderer value: %w", err) + } + err = sstore.UpdateLineRenderer(ctx, ids.ScreenId, lineId, renderer) + if err != nil { + return nil, fmt.Errorf("error changing line renderer: %v", err) + } + varsUpdated = append(varsUpdated, KwArgRenderer) + } + if view, found := pk.Kwargs[KwArgView]; found { + if err = validateRenderer(view); err != nil { + return nil, fmt.Errorf("invalid view value: %w", err) + } + err = sstore.UpdateLineRenderer(ctx, ids.ScreenId, lineId, view) + if err != nil { + return nil, fmt.Errorf("error changing line view: %v", err) + } + varsUpdated = append(varsUpdated, KwArgView) + } + if stateJson, found := pk.Kwargs[KwArgState]; found { + if len(stateJson) > sstore.MaxLineStateSize { + return nil, fmt.Errorf("invalid state value (too large), size[%d], max[%d]", len(stateJson), sstore.MaxLineStateSize) + } + var stateMap map[string]any + err = json.Unmarshal([]byte(stateJson), &stateMap) + if err != nil { + return nil, fmt.Errorf("invalid state value, cannot parse json: %v", err) + } + err = sstore.UpdateLineState(ctx, ids.ScreenId, lineId, stateMap) + if err != nil { + return nil, fmt.Errorf("cannot update linestate: %v", err) + } + varsUpdated = append(varsUpdated, KwArgState) + } + if len(varsUpdated) == 0 { + return nil, fmt.Errorf("/line:set requires a value to set: %s", formatStrs([]string{KwArgView, KwArgState}, "or", false)) + } + updatedLine, err := sstore.GetLineById(ctx, ids.ScreenId, lineId) + if err != nil { + return nil, fmt.Errorf("/line:set cannot retrieve updated line: %v", err) + } + update := &sstore.ModelUpdate{ + Line: updatedLine, + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("line updated %s", formatStrs(varsUpdated, "and", false)), + TimeoutMs: 2000, + }, + } + return update, nil +} + +func LineViewCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) != 3 { + return nil, fmt.Errorf("usage /line:view [session] [screen] [line]") + } + sessionArg := pk.Args[0] + screenArg := pk.Args[1] + lineArg := pk.Args[2] + sessionId, err := resolveSessionArg(sessionArg) + if err != nil { + return nil, fmt.Errorf("/line:view invalid session arg: %v", err) + } + if sessionId == "" { + return nil, fmt.Errorf("/line:view no session found") + } + screenRItem, err := resolveSessionScreen(ctx, sessionId, screenArg, "") + if err != nil { + return nil, fmt.Errorf("/line:view invalid screen arg: %v", err) + } + if screenRItem == nil { + return nil, fmt.Errorf("/line:view no screen found") + } + screen, err := sstore.GetScreenById(ctx, screenRItem.Id) + if err != nil { + return nil, fmt.Errorf("/line:view could not get screen: %v", err) + } + lineRItem, err := resolveLine(ctx, sessionId, screen.ScreenId, lineArg, "") + if err != nil { + return nil, fmt.Errorf("/line:view invalid line arg: %v", err) + } + update, err := sstore.SwitchScreenById(ctx, sessionId, screenRItem.Id) + if err != nil { + return nil, err + } + if lineRItem != nil { + updateMap := make(map[string]interface{}) + updateMap[sstore.ScreenField_SelectedLine] = lineRItem.Num + updateMap[sstore.ScreenField_AnchorLine] = lineRItem.Num + updateMap[sstore.ScreenField_AnchorOffset] = 0 + screen, err = sstore.UpdateScreen(ctx, screenRItem.Id, updateMap) + if err != nil { + return nil, err + } + update.Screens = []*sstore.ScreenType{screen} + } + return update, nil +} + +func BookmarksShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + // no resolve ui ids! + var tagName string // defaults to '' + if len(pk.Args) > 0 { + tagName = pk.Args[0] + } + bms, err := sstore.GetBookmarks(ctx, tagName) + if err != nil { + return nil, fmt.Errorf("cannot retrieve bookmarks: %v", err) + } + err = sstore.UpdateCurrentActivity(ctx, sstore.ActivityUpdate{BookmarksView: 1}) + if err != nil { + log.Printf("error updating current activity (bookmarks): %v\n", err) + } + update := &sstore.ModelUpdate{ + MainView: sstore.MainViewBookmarks, + Bookmarks: bms, + } + return update, nil +} + +func BookmarkSetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/bookmark:set requires one argument (bookmark id)") + } + bookmarkArg := pk.Args[0] + bookmarkId, err := sstore.GetBookmarkIdByArg(ctx, bookmarkArg) + if err != nil { + return nil, fmt.Errorf("error trying to resolve bookmark: %v", err) + } + if bookmarkId == "" { + return nil, fmt.Errorf("bookmark not found") + } + editMap := make(map[string]interface{}) + if descStr, found := pk.Kwargs["desc"]; found { + editMap[sstore.BookmarkField_Desc] = descStr + } + if cmdStr, found := pk.Kwargs["cmdstr"]; found { + editMap[sstore.BookmarkField_CmdStr] = cmdStr + } + if len(editMap) == 0 { + return nil, fmt.Errorf("no fields set, can set %s", formatStrs([]string{"desc", "cmdstr"}, "or", false)) + } + err = sstore.EditBookmark(ctx, bookmarkId, editMap) + if err != nil { + return nil, fmt.Errorf("error trying to edit bookmark: %v", err) + } + bm, err := sstore.GetBookmarkById(ctx, bookmarkId, "") + if err != nil { + return nil, fmt.Errorf("error retrieving edited bookmark: %v", err) + } + return &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoMsg: "bookmark edited", + }, + Bookmarks: []*sstore.BookmarkType{bm}, + }, nil +} + +func BookmarkDeleteCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/bookmark:delete requires one argument (bookmark id)") + } + bookmarkArg := pk.Args[0] + bookmarkId, err := sstore.GetBookmarkIdByArg(ctx, bookmarkArg) + if err != nil { + return nil, fmt.Errorf("error trying to resolve bookmark: %v", err) + } + if bookmarkId == "" { + return nil, fmt.Errorf("bookmark not found") + } + err = sstore.DeleteBookmark(ctx, bookmarkId) + if err != nil { + return nil, fmt.Errorf("error deleting bookmark: %v", err) + } + bm := &sstore.BookmarkType{BookmarkId: bookmarkId, Remove: true} + return &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoMsg: "bookmark deleted", + }, + Bookmarks: []*sstore.BookmarkType{bm}, + }, nil +} + +func LineBookmarkCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/line:bookmark requires an argument (line number or id)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + if lineId == "" { + return nil, fmt.Errorf("line %q not found", lineArg) + } + _, cmdObj, err := sstore.GetLineCmdByLineId(ctx, ids.ScreenId, lineId) + if err != nil { + return nil, fmt.Errorf("/line:bookmark error getting line: %v", err) + } + if cmdObj == nil { + return nil, fmt.Errorf("cannot bookmark non-cmd line") + } + existingBmIds, err := sstore.GetBookmarkIdsByCmdStr(ctx, cmdObj.CmdStr) + if err != nil { + return nil, fmt.Errorf("error trying to retrieve current boookmarks: %v", err) + } + var newBmId string + if len(existingBmIds) > 0 { + newBmId = existingBmIds[0] + } else { + newBm := &sstore.BookmarkType{ + BookmarkId: uuid.New().String(), + CreatedTs: time.Now().UnixMilli(), + CmdStr: cmdObj.CmdStr, + Alias: "", + Tags: nil, + Description: "", + } + err = sstore.InsertBookmark(ctx, newBm) + if err != nil { + return nil, fmt.Errorf("cannot insert bookmark: %v", err) + } + newBmId = newBm.BookmarkId + } + bms, err := sstore.GetBookmarks(ctx, "") + update := &sstore.ModelUpdate{ + MainView: sstore.MainViewBookmarks, + Bookmarks: bms, + SelectedBookmark: newBmId, + } + return update, nil +} + +func LinePinCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + return nil, nil +} + +func LineStarCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/line:star requires an argument (line number or id)") + } + if len(pk.Args) > 2 { + return nil, fmt.Errorf("/line:star only takes up to 2 arguments (line-number and star-value)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + if lineId == "" { + return nil, fmt.Errorf("line %q not found", lineArg) + } + starVal, err := resolveNonNegInt(pk.Args[1], 1) + if err != nil { + return nil, fmt.Errorf("/line:star invalid star-value (not integer): %v", err) + } + if starVal > 5 { + return nil, fmt.Errorf("/line:star invalid star-value must be in the range of 0-5") + } + err = sstore.UpdateLineStar(ctx, ids.ScreenId, lineId, starVal) + if err != nil { + return nil, fmt.Errorf("/line:star error updating star value: %v", err) + } + lineObj, err := sstore.GetLineById(ctx, ids.ScreenId, lineId) + if err != nil { + return nil, fmt.Errorf("/line:star error getting line: %v", err) + } + if lineObj == nil { + // no line (which is strange given we checked for it above). just return a nop. + return nil, nil + } + return &sstore.ModelUpdate{Line: lineObj}, nil +} + +func LineArchiveCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/line:archive requires an argument (line number or id)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + if lineId == "" { + return nil, fmt.Errorf("line %q not found", lineArg) + } + shouldArchive := true + if len(pk.Args) >= 2 { + shouldArchive = resolveBool(pk.Args[1], true) + } + err = sstore.SetLineArchivedById(ctx, ids.ScreenId, lineId, shouldArchive) + if err != nil { + return nil, fmt.Errorf("/line:archive error updating hidden status: %v", err) + } + lineObj, err := sstore.GetLineById(ctx, ids.ScreenId, lineId) + if err != nil { + return nil, fmt.Errorf("/line:archive error getting line: %v", err) + } + if lineObj == nil { + // no line (which is strange given we checked for it above). just return a nop. + return nil, nil + } + return &sstore.ModelUpdate{Line: lineObj}, nil +} + +func LinePurgeCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/line:purge requires at least one argument (line number or id)") + } + var lineIds []string + for _, lineArg := range pk.Args { + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + if lineId == "" { + return nil, fmt.Errorf("line %q not found", lineArg) + } + lineIds = append(lineIds, lineId) + } + err = sstore.PurgeLinesByIds(ctx, ids.ScreenId, lineIds) + if err != nil { + return nil, fmt.Errorf("/line:purge error purging lines: %v", err) + } + update := &sstore.ModelUpdate{} + for _, lineId := range lineIds { + lineObj := &sstore.LineType{ + ScreenId: ids.ScreenId, + LineId: lineId, + Remove: true, + } + update.Lines = append(update.Lines, lineObj) + } + return update, nil +} + +func LineShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/line:show requires an argument (line number or id)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + if lineId == "" { + return nil, fmt.Errorf("line %q not found", lineArg) + } + line, cmd, err := sstore.GetLineCmdByLineId(ctx, ids.ScreenId, lineId) + if err != nil { + return nil, fmt.Errorf("error getting line: %v", err) + } + if line == nil { + return nil, fmt.Errorf("line %q not found", lineArg) + } + var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "screenid", line.ScreenId)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "lineid", line.LineId)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "type", line.LineType)) + lineNumStr := strconv.FormatInt(line.LineNum, 10) + if line.LineNumTemp { + lineNumStr = "~" + lineNumStr + } + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "linenum", lineNumStr)) + ts := time.UnixMilli(line.Ts) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "ts", ts.Format(TsFormatStr))) + if line.Ephemeral { + buf.WriteString(fmt.Sprintf(" %-15s %v\n", "ephemeral", true)) + } + if line.Renderer != "" { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "renderer", line.Renderer)) + } else { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "renderer", "terminal")) + } + if cmd != nil { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "remote", cmd.Remote.MakeFullRemoteRef())) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "status", cmd.Status)) + if cmd.FeState["cwd"] != "" { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "cwd", cmd.FeState["cwd"])) + } + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "termopts", formatTermOpts(cmd.TermOpts))) + if cmd.TermOpts != cmd.OrigTermOpts { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "orig-termopts", formatTermOpts(cmd.OrigTermOpts))) + } + if cmd.RtnState { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "rtnstate", "true")) + } + stat, _ := sstore.StatCmdPtyFile(ctx, cmd.ScreenId, cmd.LineId) + if stat == nil { + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "file", "-")) + } else { + fileDataStr := fmt.Sprintf("v%d data=%d offset=%d max=%s", stat.Version, stat.DataSize, stat.FileOffset, scbase.NumFormatB2(stat.MaxSize)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "file", stat.Location)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "file-data", fileDataStr)) + } + if cmd.DoneTs != 0 { + doneTs := time.UnixMilli(cmd.DoneTs) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "donets", doneTs.Format(TsFormatStr))) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "exitcode", cmd.ExitCode)) + buf.WriteString(fmt.Sprintf(" %-15s %dms\n", "duration", cmd.DurationMs)) + } + } + stateStr := dbutil.QuickJson(line.LineState) + if len(stateStr) > 80 { + stateStr = stateStr[0:77] + "..." + } + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "state", stateStr)) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("line %d info", line.LineNum), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func SetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + var setMap map[string]map[string]string + setMap = make(map[string]map[string]string) + _, err := resolveUiIds(ctx, pk, 0) // best effort + if err != nil { + return nil, err + } + for argIdx, rawArgVal := range pk.Args { + eqIdx := strings.Index(rawArgVal, "=") + if eqIdx == -1 { + return nil, fmt.Errorf("/set invalid argument %d, does not contain an '='", argIdx) + } + argName := rawArgVal[:eqIdx] + argVal := rawArgVal[eqIdx+1:] + ok, scopeName, varName := resolveSetArg(argName) + if !ok { + return nil, fmt.Errorf("/set invalid setvar %q", argName) + } + if _, ok := setMap[scopeName]; !ok { + setMap[scopeName] = make(map[string]string) + } + setMap[scopeName][varName] = argVal + } + return nil, nil +} + +func makeStreamFilePk(ids resolvedIds, pk *scpacket.FeCommandPacketType) (*packet.StreamFilePacketType, error) { + cwd := ids.Remote.FeState["cwd"] + fileArg := pk.Args[0] + if fileArg == "" { + return nil, fmt.Errorf("/view:stat file argument must be set (cannot be empty)") + } + streamPk := packet.MakeStreamFilePacket() + streamPk.ReqId = uuid.New().String() + if filepath.IsAbs(fileArg) { + streamPk.Path = fileArg + } else { + streamPk.Path = filepath.Join(cwd, fileArg) + } + return streamPk, nil +} + +func ViewStatCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/view:stat requires an argument (file name)") + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + streamPk, err := makeStreamFilePk(ids, pk) + if err != nil { + return nil, err + } + streamPk.StatOnly = true + msh := ids.Remote.MShell + iter, err := msh.StreamFile(ctx, streamPk) + if err != nil { + return nil, fmt.Errorf("/view:stat error: %v", err) + } + defer iter.Close() + respIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/view:stat error getting response: %v", err) + } + resp, ok := respIf.(*packet.StreamFileResponseType) + if !ok { + return nil, fmt.Errorf("/view:stat error, bad response packet type: %T", respIf) + } + if resp.Error != "" { + return nil, fmt.Errorf("/view:stat error: %s", resp.Error) + } + if resp.Info == nil { + return nil, fmt.Errorf("/view:stat error, no file info") + } + var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "path", resp.Info.Name)) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "size", resp.Info.Size)) + modTs := time.UnixMilli(resp.Info.ModTs) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "modts", modTs.Format(TsFormatStr))) + buf.WriteString(fmt.Sprintf(" %-15s %v\n", "isdir", resp.Info.IsDir)) + modeStr := fs.FileMode(resp.Info.Perm).String() + if len(modeStr) > 9 { + modeStr = modeStr[len(modeStr)-9:] + } + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "perms", modeStr)) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("view stat %q", streamPk.Path), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func ViewTestCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/view:test requires an argument (file name)") + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + streamPk, err := makeStreamFilePk(ids, pk) + if err != nil { + return nil, err + } + msh := ids.Remote.MShell + iter, err := msh.StreamFile(ctx, streamPk) + if err != nil { + return nil, fmt.Errorf("/view:test error: %v", err) + } + defer iter.Close() + respIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/view:test error getting response: %v", err) + } + resp, ok := respIf.(*packet.StreamFileResponseType) + if !ok { + return nil, fmt.Errorf("/view:test error, bad response packet type: %T", respIf) + } + if resp.Error != "" { + return nil, fmt.Errorf("/view:test error: %s", resp.Error) + } + if resp.Info == nil { + return nil, fmt.Errorf("/view:test error, no file info") + } + var buf bytes.Buffer + var numPackets int + for { + dataPkIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/view:test error while getting data: %w", err) + } + if dataPkIf == nil { + break + } + dataPk, ok := dataPkIf.(*packet.FileDataPacketType) + if !ok { + return nil, fmt.Errorf("/view:test invalid data packet type: %T", dataPkIf) + } + if dataPk.Error != "" { + return nil, fmt.Errorf("/view:test error returned while getting data: %s", dataPk.Error) + } + numPackets++ + buf.Write(dataPk.Data) + } + buf.WriteString(fmt.Sprintf("\n\ntotal packets: %d\n", numPackets)) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("view file %q", streamPk.Path), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func CodeEditCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("%s requires an argument (file name)", GetCmdStr(pk)) + } + // TODO more error checking on filename format? + if pk.Args[0] == "" { + return nil, fmt.Errorf("%s argument cannot be empty", GetCmdStr(pk)) + } + langArg, err := getLangArg(pk) + if err != nil { + return nil, fmt.Errorf("%s invalid 'lang': %v", GetCmdStr(pk), err) + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + outputStr := fmt.Sprintf("%s %q", GetCmdStr(pk), pk.Args[0]) + cmd, err := makeStaticCmd(ctx, GetCmdStr(pk), ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + // set the line state + lineState := make(map[string]any) + lineState[sstore.LineState_Source] = "file" + lineState[sstore.LineState_File] = pk.Args[0] + if GetCmdStr(pk) == "codeview" { + lineState[sstore.LineState_Mode] = "view" + } else { + lineState[sstore.LineState_Mode] = "edit" + } + if langArg != "" { + lineState[sstore.LineState_Lang] = langArg + } + update, err := addLineForCmd(ctx, "/"+GetCmdStr(pk), true, ids, cmd, "code", lineState) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + return update, nil +} + +func CSVViewCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("%s requires an argument (file name)", GetCmdStr(pk)) + } + // TODO more error checking on filename format? + if pk.Args[0] == "" { + return nil, fmt.Errorf("%s argument cannot be empty", GetCmdStr(pk)) + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + outputStr := fmt.Sprintf("%s %q", GetCmdStr(pk), pk.Args[0]) + cmd, err := makeStaticCmd(ctx, GetCmdStr(pk), ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + // set the line state + lineState := make(map[string]any) + lineState[sstore.LineState_Source] = "file" + lineState[sstore.LineState_File] = pk.Args[0] + update, err := addLineForCmd(ctx, "/"+GetCmdStr(pk), true, ids, cmd, "csv", lineState) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + return update, nil +} + +func ImageViewCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("%s requires an argument (file name)", GetCmdStr(pk)) + } + // TODO more error checking on filename format? + if pk.Args[0] == "" { + return nil, fmt.Errorf("%s argument cannot be empty", GetCmdStr(pk)) + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + outputStr := fmt.Sprintf("%s %q", GetCmdStr(pk), pk.Args[0]) + cmd, err := makeStaticCmd(ctx, GetCmdStr(pk), ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + // set the line state + lineState := make(map[string]any) + lineState[sstore.LineState_Source] = "file" + lineState[sstore.LineState_File] = pk.Args[0] + update, err := addLineForCmd(ctx, "/"+GetCmdStr(pk), false, ids, cmd, "image", lineState) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + return update, nil +} + +func MarkdownViewCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("%s requires an argument (file name)", GetCmdStr(pk)) + } + // TODO more error checking on filename format? + if pk.Args[0] == "" { + return nil, fmt.Errorf("%s argument cannot be empty", GetCmdStr(pk)) + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + outputStr := fmt.Sprintf("%s %q", GetCmdStr(pk), pk.Args[0]) + cmd, err := makeStaticCmd(ctx, GetCmdStr(pk), ids, pk.GetRawStr(), []byte(outputStr)) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + // set the line state + lineState := make(map[string]any) + lineState[sstore.LineState_Source] = "file" + lineState[sstore.LineState_File] = pk.Args[0] + update, err := addLineForCmd(ctx, "/"+GetCmdStr(pk), false, ids, cmd, "markdown", lineState) + if err != nil { + // TODO tricky error since the command was a success, but we can't show the output + return nil, err + } + update.Interactive = pk.Interactive + return update, nil +} + +func EditTestCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/edit:test requires an argument (file name)") + } + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen|R_RemoteConnected) + if err != nil { + return nil, err + } + content, ok := pk.Kwargs["content"] + if !ok { + return nil, fmt.Errorf("/edit:test no content for file specified") + } + fileArg := pk.Args[0] + if fileArg == "" { + return nil, fmt.Errorf("/view:stat file argument must be set (cannot be empty)") + } + writePk := packet.MakeWriteFilePacket() + writePk.ReqId = uuid.New().String() + writePk.UseTemp = true + cwd := ids.Remote.FeState["cwd"] + if filepath.IsAbs(fileArg) { + writePk.Path = fileArg + } else { + writePk.Path = filepath.Join(cwd, fileArg) + } + msh := ids.Remote.MShell + iter, err := msh.PacketRpcIter(ctx, writePk) + if err != nil { + return nil, fmt.Errorf("/edit:test error: %v", err) + } + // first packet should be WriteFileReady + readyIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/edit:test error while getting ready response: %w", err) + } + readyPk, ok := readyIf.(*packet.WriteFileReadyPacketType) + if !ok { + return nil, fmt.Errorf("/edit:test bad ready packet received: %T", readyIf) + } + if readyPk.Error != "" { + return nil, fmt.Errorf("/edit:test %s", readyPk.Error) + } + dataPk := packet.MakeFileDataPacket(writePk.ReqId) + dataPk.Data = []byte(content) + dataPk.Eof = true + err = msh.SendFileData(dataPk) + if err != nil { + return nil, fmt.Errorf("/edit:test error sending data packet: %v", err) + } + doneIf, err := iter.Next(ctx) + if err != nil { + return nil, fmt.Errorf("/edit:test error while getting done response: %w", err) + } + donePk, ok := doneIf.(*packet.WriteFileDonePacketType) + if !ok { + return nil, fmt.Errorf("/edit:test bad done packet received: %T", doneIf) + } + if donePk.Error != "" { + return nil, fmt.Errorf("/edit:test %s", donePk.Error) + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("edit test, wrote %q", writePk.Path), + }, + } + return update, nil +} + +func SignalCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + ids, err := resolveUiIds(ctx, pk, R_Session|R_Screen) + if err != nil { + return nil, err + } + if len(pk.Args) == 0 { + return nil, fmt.Errorf("/signal requires a first argument (line number or id)") + } + if len(pk.Args) == 1 { + return nil, fmt.Errorf("/signal requires a second argument (signal name)") + } + lineArg := pk.Args[0] + lineId, err := sstore.FindLineIdByArg(ctx, ids.ScreenId, lineArg) + if err != nil { + return nil, fmt.Errorf("error looking up lineid: %v", err) + } + line, cmd, err := sstore.GetLineCmdByLineId(ctx, ids.ScreenId, lineId) + if err != nil { + return nil, fmt.Errorf("error getting line: %v", err) + } + if line == nil { + return nil, fmt.Errorf("line %q not found", lineArg) + } + if cmd == nil { + return nil, fmt.Errorf("line %q does not have a command", lineArg) + } + if cmd.Status != sstore.CmdStatusRunning { + return nil, fmt.Errorf("line %q command is not running, cannot send signal", lineArg) + } + sigArg := pk.Args[1] + if isAllDigits(sigArg) { + val, _ := strconv.Atoi(sigArg) + if val <= 0 || val > MaxSignalNum { + return nil, fmt.Errorf("signal number is out of bounds: %q", sigArg) + } + } else if !strings.HasPrefix(sigArg, "SIG") { + sigArg = "SIG" + sigArg + } + sigArg = strings.ToUpper(sigArg) + if len(sigArg) > 12 { + return nil, fmt.Errorf("invalid signal (too long): %q", sigArg) + } + if !sigNameRe.MatchString(sigArg) { + return nil, fmt.Errorf("invalid signal name/number: %q", sigArg) + } + msh := remote.GetRemoteById(cmd.Remote.RemoteId) + if msh == nil { + return nil, fmt.Errorf("cannot send signal, no remote found for command") + } + if !msh.IsConnected() { + return nil, fmt.Errorf("cannot send signal, remote is not connected") + } + siPk := packet.MakeSpecialInputPacket() + siPk.CK = base.MakeCommandKey(cmd.ScreenId, cmd.LineId) + siPk.SigName = sigArg + err = msh.SendSpecialInput(siPk) + if err != nil { + return nil, fmt.Errorf("cannot send signal: %v", err) + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("sent line %s signal %s", lineArg, sigArg), + }, + } + return update, nil +} + +func KillServerCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + go func() { + log.Printf("received /killserver, shutting down\n") + time.Sleep(1 * time.Second) + syscall.Kill(syscall.Getpid(), syscall.SIGINT) + }() + return nil, nil +} + +func ClientCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + return nil, fmt.Errorf("/client requires a subcommand: %s", formatStrs([]string{"show", "set"}, "or", false)) +} + +func ClientNotifyUpdateWriterCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + pcloud.ResetUpdateWriterNumFailures() + sstore.NotifyUpdateWriter() + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("notified update writer"), + }, + } + return update, nil +} + +func boolToStr(v bool, trueStr string, falseStr string) string { + if v { + return trueStr + } + return falseStr +} + +func ClientAcceptTosCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + clientOpts := clientData.ClientOpts + clientOpts.AcceptedTos = time.Now().UnixMilli() + err = sstore.SetClientOpts(ctx, clientOpts) + if err != nil { + return nil, fmt.Errorf("error updating client data: %v", err) + } + clientData, err = sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve updated client data: %v", err) + } + update := &sstore.ModelUpdate{ + ClientData: clientData, + } + return update, nil +} + +func validateOpenAIAPIToken(key string) error { + if len(key) == 0 { + return fmt.Errorf("invalid openai token, zero length") + } + if len(key) > MaxOpenAIAPITokenLen { + return fmt.Errorf("invalid openai token, too long") + } + for idx, ch := range key { + if !unicode.IsPrint(ch) { + return fmt.Errorf("invalid openai token, char at idx:%d is invalid %q", idx, string(ch)) + } + } + return nil +} + +func validateOpenAIModel(model string) error { + if len(model) == 0 { + return nil + } + if len(model) > MaxOpenAIModelLen { + return fmt.Errorf("invalid openai model, too long") + } + for idx, ch := range model { + if !unicode.IsPrint(ch) { + return fmt.Errorf("invalid openai model, char at idx:%d is invalid %q", idx, string(ch)) + } + } + return nil +} + +func ClientSetCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + var varsUpdated []string + if fontSizeStr, found := pk.Kwargs["termfontsize"]; found { + newFontSize, err := resolveNonNegInt(fontSizeStr, 0) + if err != nil { + return nil, fmt.Errorf("invalid termfontsize, must be a number between 8-15: %v", err) + } + if newFontSize < 8 || newFontSize > 15 { + return nil, fmt.Errorf("invalid termfontsize, must be a number between 8-15") + } + feOpts := clientData.FeOpts + feOpts.TermFontSize = newFontSize + err = sstore.UpdateClientFeOpts(ctx, feOpts) + if err != nil { + return nil, fmt.Errorf("error updating client feopts: %v", err) + } + varsUpdated = append(varsUpdated, "termfontsize") + } + if apiToken, found := pk.Kwargs["openaiapitoken"]; found { + err = validateOpenAIAPIToken(apiToken) + if err != nil { + return nil, err + } + varsUpdated = append(varsUpdated, "openaiapitoken") + aiOpts := clientData.OpenAIOpts + if aiOpts == nil { + aiOpts = &sstore.OpenAIOptsType{} + clientData.OpenAIOpts = aiOpts + } + aiOpts.APIToken = apiToken + err = sstore.UpdateClientOpenAIOpts(ctx, *aiOpts) + if err != nil { + return nil, fmt.Errorf("error updating client openai api token: %v", err) + } + } + if aiModel, found := pk.Kwargs["openaimodel"]; found { + err = validateOpenAIModel(aiModel) + if err != nil { + return nil, err + } + varsUpdated = append(varsUpdated, "openaimodel") + aiOpts := clientData.OpenAIOpts + if aiOpts == nil { + aiOpts = &sstore.OpenAIOptsType{} + clientData.OpenAIOpts = aiOpts + } + aiOpts.Model = aiModel + err = sstore.UpdateClientOpenAIOpts(ctx, *aiOpts) + if err != nil { + return nil, fmt.Errorf("error updating client openai model: %v", err) + } + } + if maxTokensStr, found := pk.Kwargs["openaimaxtokens"]; found { + maxTokens, err := strconv.Atoi(maxTokensStr) + if err != nil { + return nil, fmt.Errorf("error updating client openai maxtokens, invalid number: %v", err) + } + if maxTokens < 0 || maxTokens > 1000000 { + return nil, fmt.Errorf("error updating client openai maxtokens, out of range: %d", maxTokens) + } + varsUpdated = append(varsUpdated, "openaimaxtokens") + aiOpts := clientData.OpenAIOpts + if aiOpts == nil { + aiOpts = &sstore.OpenAIOptsType{} + clientData.OpenAIOpts = aiOpts + } + aiOpts.MaxTokens = maxTokens + err = sstore.UpdateClientOpenAIOpts(ctx, *aiOpts) + if err != nil { + return nil, fmt.Errorf("error updating client openai maxtokens: %v", err) + } + } + if maxChoicesStr, found := pk.Kwargs["openaimaxchoices"]; found { + maxChoices, err := strconv.Atoi(maxChoicesStr) + if err != nil { + return nil, fmt.Errorf("error updating client openai maxchoices, invalid number: %v", err) + } + if maxChoices < 0 || maxChoices > 10 { + return nil, fmt.Errorf("error updating client openai maxchoices, out of range: %d", maxChoices) + } + varsUpdated = append(varsUpdated, "openaimaxchoices") + aiOpts := clientData.OpenAIOpts + if aiOpts == nil { + aiOpts = &sstore.OpenAIOptsType{} + clientData.OpenAIOpts = aiOpts + } + aiOpts.MaxChoices = maxChoices + err = sstore.UpdateClientOpenAIOpts(ctx, *aiOpts) + if err != nil { + return nil, fmt.Errorf("error updating client openai maxchoices: %v", err) + } + } + if len(varsUpdated) == 0 { + return nil, fmt.Errorf("/client:set requires a value to set: %s", formatStrs([]string{"termfontsize", "openaiapitoken", "openaimodel", "openaimaxtokens", "openaimaxchoices"}, "or", false)) + } + clientData, err = sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve updated client data: %v", err) + } + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoMsg: fmt.Sprintf("client updated %s", formatStrs(varsUpdated, "and", false)), + TimeoutMs: 2000, + }, + ClientData: clientData, + } + return update, nil +} + +func ClientShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + dbVersion, err := sstore.GetDBVersion(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve db version: %v\n", err) + } + clientVersion := "-" + if pk.UIContext != nil && pk.UIContext.Build != "" { + clientVersion = pk.UIContext.Build + } + var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "userid", clientData.UserId)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "clientid", clientData.ClientId)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "telemetry", boolToStr(clientData.ClientOpts.NoTelemetry, "off", "on"))) + buf.WriteString(fmt.Sprintf(" %-15s %d\n", "db-version", dbVersion)) + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "client-version", clientVersion)) + buf.WriteString(fmt.Sprintf(" %-15s %s %s\n", "server-version", scbase.PromptVersion, scbase.BuildTime)) + buf.WriteString(fmt.Sprintf(" %-15s %s (%s)\n", "arch", scbase.ClientArch(), scbase.MacOSRelease())) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("client info"), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func TelemetryCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + return nil, fmt.Errorf("/telemetry requires a subcommand: %s", formatStrs([]string{"show", "on", "off", "send"}, "or", false)) +} + +func setNoTelemetry(ctx context.Context, clientData *sstore.ClientData, noTelemetryVal bool) error { + clientOpts := clientData.ClientOpts + clientOpts.NoTelemetry = noTelemetryVal + err := sstore.SetClientOpts(ctx, clientOpts) + if err != nil { + return fmt.Errorf("error trying to update client telemetry: %v", err) + } + log.Printf("client no-telemetry setting updated to %v\n", noTelemetryVal) + err = pcloud.SendNoTelemetryUpdate(ctx, clientOpts.NoTelemetry) + if err != nil { + // ignore error, just log + log.Printf("[error] sending no-telemetry update: %v\n", err) + log.Printf("note that telemetry update has still taken effect locally, and will be respected by the client\n") + } + return nil +} + +func TelemetryOnCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + if !clientData.ClientOpts.NoTelemetry { + return sstore.InfoMsgUpdate("telemetry is already on"), nil + } + err = setNoTelemetry(ctx, clientData, false) + if err != nil { + return nil, err + } + err = pcloud.SendTelemetry(ctx, false) + if err != nil { + // ignore error, but log + log.Printf("[error] sending telemetry update (in /telemetry:on): %v\n", err) + } + clientData, err = sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve updated client data: %v", err) + } + update := sstore.InfoMsgUpdate("telemetry is now on") + update.ClientData = clientData + return update, nil +} + +func TelemetryOffCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + if clientData.ClientOpts.NoTelemetry { + return sstore.InfoMsgUpdate("telemetry is already off"), nil + } + err = setNoTelemetry(ctx, clientData, true) + if err != nil { + return nil, err + } + clientData, err = sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve updated client data: %v", err) + } + update := sstore.InfoMsgUpdate("telemetry is now off") + update.ClientData = clientData + return update, nil +} + +func TelemetryShowCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + var buf bytes.Buffer + buf.WriteString(fmt.Sprintf(" %-15s %s\n", "telemetry", boolToStr(clientData.ClientOpts.NoTelemetry, "off", "on"))) + update := &sstore.ModelUpdate{ + Info: &sstore.InfoMsgType{ + InfoTitle: fmt.Sprintf("telemetry info"), + InfoLines: splitLinesForInfo(buf.String()), + }, + } + return update, nil +} + +func TelemetrySendCommand(ctx context.Context, pk *scpacket.FeCommandPacketType) (sstore.UpdatePacket, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return nil, fmt.Errorf("cannot retrieve client data: %v", err) + } + force := resolveBool(pk.Kwargs["force"], false) + if clientData.ClientOpts.NoTelemetry && !force { + return nil, fmt.Errorf("cannot send telemetry, telemetry is off. pass force=1 to force the send, or turn on telemetry with /telemetry:on") + } + err = pcloud.SendTelemetry(ctx, force) + if err != nil { + return nil, fmt.Errorf("failed to send telemetry: %v", err) + } + return sstore.InfoMsgUpdate("telemetry sent"), nil +} + +func formatTermOpts(termOpts sstore.TermOpts) string { + if termOpts.Cols == 0 { + return "???" + } + rtnStr := fmt.Sprintf("%dx%d", termOpts.Rows, termOpts.Cols) + if termOpts.FlexRows { + rtnStr += " flexrows" + } + if termOpts.MaxPtySize > 0 { + rtnStr += " maxbuf=" + scbase.NumFormatB2(termOpts.MaxPtySize) + } + return rtnStr +} + +type ColMeta struct { + Title string + MinCols int + MaxCols int +} + +func toInterfaceArr(sarr []string) []interface{} { + rtn := make([]interface{}, len(sarr)) + for idx, s := range sarr { + rtn[idx] = s + } + return rtn +} + +func formatTextTable(totalCols int, data [][]string, colMeta []ColMeta) []string { + numCols := len(colMeta) + maxColLen := make([]int, len(colMeta)) + for i, cm := range colMeta { + maxColLen[i] = cm.MinCols + } + for _, row := range data { + for i := 0; i < numCols && i < len(row); i++ { + dlen := len(row[i]) + if dlen > maxColLen[i] { + maxColLen[i] = dlen + } + } + } + fmtStr := "" + for idx, clen := range maxColLen { + if idx != 0 { + fmtStr += " " + } + fmtStr += fmt.Sprintf("%%%ds", clen) + } + var rtn []string + for _, row := range data { + sval := fmt.Sprintf(fmtStr, toInterfaceArr(row)...) + rtn = append(rtn, sval) + } + return rtn +} + +func isValidInScope(scopeName string, varName string) bool { + for _, varScope := range SetVarScopes { + if varScope.ScopeName == scopeName { + return utilfn.ContainsStr(varScope.VarNames, varName) + } + } + return false +} + +// returns (is-valid, scope, name) +// TODO write a full resolver to allow for indexed arguments. e.g. session[1].screen[1].screen.pterm="25x80" +func resolveSetArg(argName string) (bool, string, string) { + dotIdx := strings.Index(argName, ".") + if dotIdx == -1 { + argName = SetVarNameMap[argName] + dotIdx = strings.Index(argName, ".") + } + if argName == "" { + return false, "", "" + } + scopeName := argName[0:dotIdx] + varName := argName[dotIdx+1:] + if !isValidInScope(scopeName, varName) { + return false, "", "" + } + return true, scopeName, varName +} diff --git a/wavesrv/pkg/cmdrunner/linux-decls.txt b/wavesrv/pkg/cmdrunner/linux-decls.txt new file mode 100644 index 00000000..acb5c807 --- /dev/null +++ b/wavesrv/pkg/cmdrunner/linux-decls.txt @@ -0,0 +1,1986 @@ +__expand_tilde_by_ref () +{ + if [[ ${!1} == \~* ]]; then + eval $1=$(printf ~%q "${!1#\~}"); + fi +} +__get_cword_at_cursor_by_ref () +{ + local cword words=(); + __reassemble_comp_words_by_ref "$1" words cword; + local i cur index=$COMP_POINT lead=${COMP_LINE:0:$COMP_POINT}; + if [[ $index -gt 0 && ( -n $lead && -n ${lead//[[:space:]]} ) ]]; then + cur=$COMP_LINE; + for ((i = 0; i <= cword; ++i )) + do + while [[ ${#cur} -ge ${#words[i]} && "${cur:0:${#words[i]}}" != "${words[i]}" ]]; do + cur="${cur:1}"; + [[ $index -gt 0 ]] && ((index--)); + done; + if [[ $i -lt $cword ]]; then + local old_size=${#cur}; + cur="${cur#"${words[i]}"}"; + local new_size=${#cur}; + (( index -= old_size - new_size )); + fi; + done; + [[ -n $cur && ! -n ${cur//[[:space:]]} ]] && cur=; + [[ $index -lt 0 ]] && index=0; + fi; + local "$2" "$3" "$4" && _upvars -a${#words[@]} $2 "${words[@]}" -v $3 "$cword" -v $4 "${cur:0:$index}" +} +__git_eread () +{ + test -r "$1" && IFS=' +' read "$2" < "$1" +} +__git_ps1 () +{ + local exit=$?; + local pcmode=no; + local detached=no; + local ps1pc_start='\u@\h:\w '; + local ps1pc_end='\$ '; + local printf_format=' (%s)'; + case "$#" in + 2 | 3) + pcmode=yes; + ps1pc_start="$1"; + ps1pc_end="$2"; + printf_format="${3:-$printf_format}"; + PS1="$ps1pc_start$ps1pc_end" + ;; + 0 | 1) + printf_format="${1:-$printf_format}" + ;; + *) + return $exit + ;; + esac; + local ps1_expanded=yes; + [ -z "${ZSH_VERSION-}" ] || [[ -o PROMPT_SUBST ]] || ps1_expanded=no; + [ -z "${BASH_VERSION-}" ] || shopt -q promptvars || ps1_expanded=no; + local repo_info rev_parse_exit_code; + repo_info="$(git rev-parse --git-dir --is-inside-git-dir --is-bare-repository --is-inside-work-tree --short HEAD 2>/dev/null)"; + rev_parse_exit_code="$?"; + if [ -z "$repo_info" ]; then + return $exit; + fi; + local short_sha=""; + if [ "$rev_parse_exit_code" = "0" ]; then + short_sha="${repo_info##* +}"; + repo_info="${repo_info% +*}"; + fi; + local inside_worktree="${repo_info##* +}"; + repo_info="${repo_info% +*}"; + local bare_repo="${repo_info##* +}"; + repo_info="${repo_info% +*}"; + local inside_gitdir="${repo_info##* +}"; + local g="${repo_info% +*}"; + if [ "true" = "$inside_worktree" ] && [ -n "${GIT_PS1_HIDE_IF_PWD_IGNORED-}" ] && [ "$(git config --bool bash.hideIfPwdIgnored)" != "false" ] && git check-ignore -q .; then + return $exit; + fi; + local r=""; + local b=""; + local step=""; + local total=""; + if [ -d "$g/rebase-merge" ]; then + __git_eread "$g/rebase-merge/head-name" b; + __git_eread "$g/rebase-merge/msgnum" step; + __git_eread "$g/rebase-merge/end" total; + if [ -f "$g/rebase-merge/interactive" ]; then + r="|REBASE-i"; + else + r="|REBASE-m"; + fi; + else + if [ -d "$g/rebase-apply" ]; then + __git_eread "$g/rebase-apply/next" step; + __git_eread "$g/rebase-apply/last" total; + if [ -f "$g/rebase-apply/rebasing" ]; then + __git_eread "$g/rebase-apply/head-name" b; + r="|REBASE"; + else + if [ -f "$g/rebase-apply/applying" ]; then + r="|AM"; + else + r="|AM/REBASE"; + fi; + fi; + else + if [ -f "$g/MERGE_HEAD" ]; then + r="|MERGING"; + else + if __git_sequencer_status; then + :; + else + if [ -f "$g/BISECT_LOG" ]; then + r="|BISECTING"; + fi; + fi; + fi; + fi; + if [ -n "$b" ]; then + :; + else + if [ -h "$g/HEAD" ]; then + b="$(git symbolic-ref HEAD 2>/dev/null)"; + else + local head=""; + if ! __git_eread "$g/HEAD" head; then + return $exit; + fi; + b="${head#ref: }"; + if [ "$head" = "$b" ]; then + detached=yes; + b="$( + case "${GIT_PS1_DESCRIBE_STYLE-}" in + (contains) + git describe --contains HEAD ;; + (branch) + git describe --contains --all HEAD ;; + (tag) + git describe --tags HEAD ;; + (describe) + git describe HEAD ;; + (* | default) + git describe --tags --exact-match HEAD ;; + esac 2>/dev/null)" || b="$short_sha..."; + b="($b)"; + fi; + fi; + fi; + fi; + if [ -n "$step" ] && [ -n "$total" ]; then + r="$r $step/$total"; + fi; + local w=""; + local i=""; + local s=""; + local u=""; + local c=""; + local p=""; + if [ "true" = "$inside_gitdir" ]; then + if [ "true" = "$bare_repo" ]; then + c="BARE:"; + else + b="GIT_DIR!"; + fi; + else + if [ "true" = "$inside_worktree" ]; then + if [ -n "${GIT_PS1_SHOWDIRTYSTATE-}" ] && [ "$(git config --bool bash.showDirtyState)" != "false" ]; then + git diff --no-ext-diff --quiet || w="*"; + git diff --no-ext-diff --cached --quiet || i="+"; + if [ -z "$short_sha" ] && [ -z "$i" ]; then + i="#"; + fi; + fi; + if [ -n "${GIT_PS1_SHOWSTASHSTATE-}" ] && git rev-parse --verify --quiet refs/stash > /dev/null; then + s="$"; + fi; + if [ -n "${GIT_PS1_SHOWUNTRACKEDFILES-}" ] && [ "$(git config --bool bash.showUntrackedFiles)" != "false" ] && git ls-files --others --exclude-standard --directory --no-empty-directory --error-unmatch -- ':/*' > /dev/null 2> /dev/null; then + u="%${ZSH_VERSION+%}"; + fi; + if [ -n "${GIT_PS1_SHOWUPSTREAM-}" ]; then + __git_ps1_show_upstream; + fi; + fi; + fi; + local z="${GIT_PS1_STATESEPARATOR-" "}"; + if [ $pcmode = yes ] && [ -n "${GIT_PS1_SHOWCOLORHINTS-}" ]; then + __git_ps1_colorize_gitstring; + fi; + b=${b##refs/heads/}; + if [ $pcmode = yes ] && [ $ps1_expanded = yes ]; then + __git_ps1_branch_name=$b; + b="\${__git_ps1_branch_name}"; + fi; + local f="$w$i$s$u"; + local gitstring="$c$b${f:+$z$f}$r$p"; + if [ $pcmode = yes ]; then + if [ "${__git_printf_supports_v-}" != yes ]; then + gitstring=$(printf -- "$printf_format" "$gitstring"); + else + printf -v gitstring -- "$printf_format" "$gitstring"; + fi; + PS1="$ps1pc_start$gitstring$ps1pc_end"; + else + printf -- "$printf_format" "$gitstring"; + fi; + return $exit +} +__git_ps1_colorize_gitstring () +{ + if [[ -n ${ZSH_VERSION-} ]]; then + local c_red='%F{red}'; + local c_green='%F{green}'; + local c_lblue='%F{blue}'; + local c_clear='%f'; + else + local c_red='\[\e[31m\]'; + local c_green='\[\e[32m\]'; + local c_lblue='\[\e[1;34m\]'; + local c_clear='\[\e[0m\]'; + fi; + local bad_color=$c_red; + local ok_color=$c_green; + local flags_color="$c_lblue"; + local branch_color=""; + if [ $detached = no ]; then + branch_color="$ok_color"; + else + branch_color="$bad_color"; + fi; + c="$branch_color$c"; + z="$c_clear$z"; + if [ "$w" = "*" ]; then + w="$bad_color$w"; + fi; + if [ -n "$i" ]; then + i="$ok_color$i"; + fi; + if [ -n "$s" ]; then + s="$flags_color$s"; + fi; + if [ -n "$u" ]; then + u="$bad_color$u"; + fi; + r="$c_clear$r" +} +__git_ps1_show_upstream () +{ + local key value; + local svn_remote svn_url_pattern count n; + local upstream=git legacy="" verbose="" name=""; + svn_remote=(); + local output="$(git config -z --get-regexp '^(svn-remote\..*\.url|bash\.showupstream)$' 2>/dev/null | tr '\0\n' '\n ')"; + while read -r key value; do + case "$key" in + bash.showupstream) + GIT_PS1_SHOWUPSTREAM="$value"; + if [[ -z "${GIT_PS1_SHOWUPSTREAM}" ]]; then + p=""; + return; + fi + ;; + svn-remote.*.url) + svn_remote[$((${#svn_remote[@]} + 1))]="$value"; + svn_url_pattern="$svn_url_pattern\\|$value"; + upstream=svn+git + ;; + esac; + done <<< "$output"; + for option in ${GIT_PS1_SHOWUPSTREAM}; + do + case "$option" in + git | svn) + upstream="$option" + ;; + verbose) + verbose=1 + ;; + legacy) + legacy=1 + ;; + name) + name=1 + ;; + esac; + done; + case "$upstream" in + git) + upstream="@{upstream}" + ;; + svn*) + local -a svn_upstream; + svn_upstream=($(git log --first-parent -1 --grep="^git-svn-id: \(${svn_url_pattern#??}\)" 2>/dev/null)); + if [[ 0 -ne ${#svn_upstream[@]} ]]; then + svn_upstream=${svn_upstream[${#svn_upstream[@]} - 2]}; + svn_upstream=${svn_upstream%@*}; + local n_stop="${#svn_remote[@]}"; + for ((n=1; n <= n_stop; n++)) + do + svn_upstream=${svn_upstream#${svn_remote[$n]}}; + done; + if [[ -z "$svn_upstream" ]]; then + upstream=${GIT_SVN_ID:-git-svn}; + else + upstream=${svn_upstream#/}; + fi; + else + if [[ "svn+git" = "$upstream" ]]; then + upstream="@{upstream}"; + fi; + fi + ;; + esac; + if [[ -z "$legacy" ]]; then + count="$(git rev-list --count --left-right "$upstream"...HEAD 2>/dev/null)"; + else + local commits; + if commits="$(git rev-list --left-right "$upstream"...HEAD 2>/dev/null)"; then + local commit behind=0 ahead=0; + for commit in $commits; + do + case "$commit" in + "<"*) + ((behind++)) + ;; + *) + ((ahead++)) + ;; + esac; + done; + count="$behind $ahead"; + else + count=""; + fi; + fi; + if [[ -z "$verbose" ]]; then + case "$count" in + "") + p="" + ;; + "0 0") + p="=" + ;; + "0 "*) + p=">" + ;; + *" 0") + p="<" + ;; + *) + p="<>" + ;; + esac; + else + case "$count" in + "") + p="" + ;; + "0 0") + p=" u=" + ;; + "0 "*) + p=" u+${count#0 }" + ;; + *" 0") + p=" u-${count% 0}" + ;; + *) + p=" u+${count#* }-${count% *}" + ;; + esac; + if [[ -n "$count" && -n "$name" ]]; then + __git_ps1_upstream_name=$(git rev-parse --abbrev-ref "$upstream" 2>/dev/null); + if [ $pcmode = yes ] && [ $ps1_expanded = yes ]; then + p="$p \${__git_ps1_upstream_name}"; + else + p="$p ${__git_ps1_upstream_name}"; + unset __git_ps1_upstream_name; + fi; + fi; + fi +} +__git_sequencer_status () +{ + local todo; + if test -f "$g/CHERRY_PICK_HEAD"; then + r="|CHERRY-PICKING"; + return 0; + else + if test -f "$g/REVERT_HEAD"; then + r="|REVERTING"; + return 0; + else + if __git_eread "$g/sequencer/todo" todo; then + case "$todo" in + p[\ \ ] | pick[\ \ ]*) + r="|CHERRY-PICKING"; + return 0 + ;; + revert[\ \ ]*) + r="|REVERTING"; + return 0 + ;; + esac; + fi; + fi; + fi; + return 1 +} +__load_completion () +{ + local -a dirs=(${BASH_COMPLETION_USER_DIR:-${XDG_DATA_HOME:-$HOME/.local/share}/bash-completion}/completions); + local OIFS=$IFS IFS=: dir cmd="${1##*/}" compfile; + [[ -n $cmd ]] || return 1; + for dir in ${XDG_DATA_DIRS:-/usr/local/share:/usr/share}; + do + dirs+=($dir/bash-completion/completions); + done; + IFS=$OIFS; + if [[ $BASH_SOURCE == */* ]]; then + dirs+=("${BASH_SOURCE%/*}/completions"); + else + dirs+=(./completions); + fi; + for dir in "${dirs[@]}"; + do + [[ -d "$dir" ]] || continue; + for compfile in "$cmd" "$cmd.bash" "_$cmd"; + do + compfile="$dir/$compfile"; + [[ -f "$compfile" ]] && . "$compfile" &> /dev/null && return 0; + done; + done; + [[ -n "${_xspecs[$cmd]}" ]] && complete -F _filedir_xspec "$cmd" && return 0; + return 1 +} +__ltrim_colon_completions () +{ + if [[ "$1" == *:* && "$COMP_WORDBREAKS" == *:* ]]; then + local colon_word=${1%"${1##*:}"}; + local i=${#COMPREPLY[*]}; + while [[ $((--i)) -ge 0 ]]; do + COMPREPLY[$i]=${COMPREPLY[$i]#"$colon_word"}; + done; + fi +} +__parse_options () +{ + local option option2 i IFS=' +,/|'; + option=; + local -a array=($1); + for i in "${array[@]}"; + do + case "$i" in + ---*) + break + ;; + --?*) + option=$i; + break + ;; + -?*) + [[ -n $option ]] || option=$i + ;; + *) + break + ;; + esac; + done; + [[ -n $option ]] || return 0; + IFS=' +'; + if [[ $option =~ (\[((no|dont)-?)\]). ]]; then + option2=${option/"${BASH_REMATCH[1]}"/}; + option2=${option2%%[<{().[]*}; + printf '%s\n' "${option2/=*/=}"; + option=${option/"${BASH_REMATCH[1]}"/"${BASH_REMATCH[2]}"}; + fi; + option=${option%%[<{().[]*}; + printf '%s\n' "${option/=*/=}" +} +__reassemble_comp_words_by_ref () +{ + local exclude i j line ref; + if [[ -n $1 ]]; then + exclude="${1//[^$COMP_WORDBREAKS]}"; + fi; + printf -v "$3" %s "$COMP_CWORD"; + if [[ -n $exclude ]]; then + line=$COMP_LINE; + for ((i=0, j=0; i < ${#COMP_WORDS[@]}; i++, j++)) + do + while [[ $i -gt 0 && ${COMP_WORDS[$i]} == +([$exclude]) ]]; do + [[ $line != [[:blank:]]* ]] && (( j >= 2 )) && ((j--)); + ref="$2[$j]"; + printf -v "$ref" %s "${!ref}${COMP_WORDS[i]}"; + [[ $i == $COMP_CWORD ]] && printf -v "$3" %s "$j"; + line=${line#*"${COMP_WORDS[$i]}"}; + [[ $line == [[:blank:]]* ]] && ((j++)); + (( $i < ${#COMP_WORDS[@]} - 1)) && ((i++)) || break 2; + done; + ref="$2[$j]"; + printf -v "$ref" %s "${!ref}${COMP_WORDS[i]}"; + line=${line#*"${COMP_WORDS[i]}"}; + [[ $i == $COMP_CWORD ]] && printf -v "$3" %s "$j"; + done; + [[ $i == $COMP_CWORD ]] && printf -v "$3" %s "$j"; + else + for i in "${!COMP_WORDS[@]}"; + do + printf -v "$2[i]" %s "${COMP_WORDS[i]}"; + done; + fi +} +_allowed_groups () +{ + if _complete_as_root; then + local IFS=' +'; + COMPREPLY=($(compgen -g -- "$1")); + else + local IFS=' + '; + COMPREPLY=($(compgen -W "$(id -Gn 2>/dev/null || groups 2>/dev/null)" -- "$1")); + fi +} +_allowed_users () +{ + if _complete_as_root; then + local IFS=' +'; + COMPREPLY=($(compgen -u -- "${1:-$cur}")); + else + local IFS=' + '; + COMPREPLY=($(compgen -W "$(id -un 2>/dev/null || whoami 2>/dev/null)" -- "${1:-$cur}")); + fi +} +_apport-bug () +{ + local cur dashoptions prev param; + COMPREPLY=(); + cur=`_get_cword`; + prev=${COMP_WORDS[COMP_CWORD-1]}; + dashoptions='-h --help --save -v --version --tag -w --window'; + case "$prev" in + ubuntu-bug | apport-bug) + case "$cur" in + -*) + COMPREPLY=($( compgen -W "$dashoptions" -- $cur )) + ;; + *) + _apport_parameterless + ;; + esac + ;; + --save) + COMPREPLY=($( compgen -o default -G "$cur*" )) + ;; + -w | --window) + dashoptions="--save --tag"; + COMPREPLY=($( compgen -W "$dashoptions" -- $cur )) + ;; + -h | --help | -v | --version | --tag) + return 0 + ;; + *) + dashoptions="--tag"; + if ! [[ "${COMP_WORDS[*]}" =~ .*--save.* ]]; then + dashoptions="--save $dashoptions"; + fi; + if ! [[ "${COMP_WORDS[*]}" =~ .*--window.* || "${COMP_WORDS[*]}" =~ .*\ -w\ .* ]]; then + dashoptions="-w --window $dashoptions"; + fi; + case "$cur" in + -*) + COMPREPLY=($( compgen -W "$dashoptions" -- $cur )) + ;; + *) + _apport_parameterless + ;; + esac + ;; + esac +} +_apport-cli () +{ + local cur dashoptions prev param; + COMPREPLY=(); + cur=`_get_cword`; + prev=${COMP_WORDS[COMP_CWORD-1]}; + dashoptions='-h --help -f --file-bug -u --update-bug -s --symptom \ + -c --crash-file --save -v --version --tag -w --window'; + case "$prev" in + apport-cli) + case "$cur" in + -*) + COMPREPLY=($( compgen -W "$dashoptions" -- $cur )) + ;; + *) + _apport_parameterless + ;; + esac + ;; + -f | --file-bug) + param="-P --pid -p --package -s --symptom"; + COMPREPLY=($( compgen -W "$param $(_apport_symptoms)" -- $cur)) + ;; + -s | --symptom) + COMPREPLY=($( compgen -W "$(_apport_symptoms)" -- $cur)) + ;; + --save) + COMPREPLY=($( compgen -o default -G "$cur*" )) + ;; + -c | --crash-file) + COMPREPLY=($( compgen -G "${cur}*.apport" + compgen -G "${cur}*.crash" )) + ;; + -w | --window) + dashoptions="--save --tag"; + COMPREPLY=($( compgen -W "$dashoptions" -- $cur )) + ;; + -h | --help | -v | --version | --tag) + return 0 + ;; + *) + dashoptions='--tag'; + if ! [[ "${COMP_WORDS[*]}" =~ .*--save.* ]]; then + dashoptions="--save $dashoptions"; + fi; + if ! [[ "${COMP_WORDS[*]}" =~ .*--window.* || "${COMP_WORDS[*]}" =~ .*\ -w\ .* ]]; then + dashoptions="-w --window $dashoptions"; + fi; + if ! [[ "${COMP_WORDS[*]}" =~ .*--symptom.* || "${COMP_WORDS[*]}" =~ .*\ -s\ .* ]]; then + dashoptions="-s --symptom $dashoptions"; + fi; + if ! [[ "${COMP_WORDS[*]}" =~ .*--update.* || "${COMP_WORDS[*]}" =~ .*\ -u\ .* ]]; then + dashoptions="-u --update $dashoptions"; + fi; + if ! [[ "${COMP_WORDS[*]}" =~ .*--file-bug.* || "${COMP_WORDS[*]}" =~ .*\ -f\ .* ]]; then + dashoptions="-f --file-bug $dashoptions"; + fi; + if ! [[ "${COMP_WORDS[*]}" =~ .*--crash-file.* || "${COMP_WORDS[*]}" =~ .*\ -c\ .* ]]; then + dashoptions="-c --crash-file $dashoptions"; + fi; + case "$cur" in + -*) + COMPREPLY=($( compgen -W "$dashoptions" -- $cur )) + ;; + *) + _apport_parameterless + ;; + esac + ;; + esac +} +_apport-collect () +{ + local cur prev; + COMPREPLY=(); + cur=`_get_cword`; + prev=${COMP_WORDS[COMP_CWORD-1]}; + case "$prev" in + apport-collect) + COMPREPLY=($( compgen -W "-p --package --tag" -- $cur)) + ;; + -p | --package) + COMPREPLY=($( apt-cache pkgnames $cur 2> /dev/null )) + ;; + --tag) + return 0 + ;; + *) + if [[ "${COMP_WORDS[*]}" =~ .*\ -p.* || "${COMP_WORDS[*]}" =~ .*--package.* ]]; then + COMPREPLY=($( compgen -W "--tag" -- $cur)); + else + COMPREPLY=($( compgen -W "-p --package --tag" -- $cur)); + fi + ;; + esac +} +_apport-unpack () +{ + local cur prev; + COMPREPLY=(); + cur=`_get_cword`; + prev=${COMP_WORDS[COMP_CWORD-1]}; + case "$prev" in + apport-unpack) + COMPREPLY=($( compgen -G "${cur}*.apport" + compgen -G "${cur}*.crash" )) + ;; + esac +} +_apport_parameterless () +{ + local param; + param="$dashoptions $( apt-cache pkgnames $cur 2> /dev/null ) $( command ps axo pid | sed 1d ) $( _apport_symptoms ) $( compgen -G "${cur}*" )"; + COMPREPLY=($( compgen -W "$param" -- $cur)) +} +_apport_symptoms () +{ + local syms; + if [ -r /usr/share/apport/symptoms ]; then + for FILE in $(ls /usr/share/apport/symptoms); + do + if [[ ! "$FILE" =~ ^_.* && -n $(egrep "^def run\s*\(.*\):" /usr/share/apport/symptoms/$FILE) ]]; then + syms="$syms ${FILE%.py}"; + fi; + done; + fi; + echo $syms +} +_available_interfaces () +{ + local PATH=$PATH:/sbin; + COMPREPLY=($({ + if [[ ${1:-} == -w ]]; then + iwconfig + elif [[ ${1:-} == -a ]]; then + ifconfig || ip link show up + else + ifconfig -a || ip link show + fi + } 2>/dev/null | awk '/^[^ \t]/ { if ($1 ~ /^[0-9]+:/) { print $2 } else { print $1 } }')); + COMPREPLY=($(compgen -W '${COMPREPLY[@]/%[[:punct:]]/}' -- "$cur")) +} +_cd () +{ + local cur prev words cword; + _init_completion || return; + local IFS=' +' i j k; + compopt -o filenames; + if [[ -z "${CDPATH:-}" || "$cur" == ?(.)?(.)/* ]]; then + _filedir -d; + return; + fi; + local -r mark_dirs=$(_rl_enabled mark-directories && echo y); + local -r mark_symdirs=$(_rl_enabled mark-symlinked-directories && echo y); + for i in ${CDPATH//:/' +'}; + do + k="${#COMPREPLY[@]}"; + for j in $(compgen -d -- $i/$cur); + do + if [[ ( -n $mark_symdirs && -h $j || -n $mark_dirs && ! -h $j ) && ! -d ${j#$i/} ]]; then + j+="/"; + fi; + COMPREPLY[k++]=${j#$i/}; + done; + done; + _filedir -d; + if [[ ${#COMPREPLY[@]} -eq 1 ]]; then + i=${COMPREPLY[0]}; + if [[ "$i" == "$cur" && $i != "*/" ]]; then + COMPREPLY[0]="${i}/"; + fi; + fi; + return +} +_cd_devices () +{ + COMPREPLY+=($(compgen -f -d -X "!*/?([amrs])cd*" -- "${cur:-/dev/}")) +} +_command () +{ + local offset i; + offset=1; + for ((i=1; i <= COMP_CWORD; i++ )) + do + if [[ "${COMP_WORDS[i]}" != -* ]]; then + offset=$i; + break; + fi; + done; + _command_offset $offset +} +_command_offset () +{ + local word_offset=$1 i j; + for ((i=0; i < $word_offset; i++ )) + do + for ((j=0; j <= ${#COMP_LINE}; j++ )) + do + [[ "$COMP_LINE" == "${COMP_WORDS[i]}"* ]] && break; + COMP_LINE=${COMP_LINE:1}; + ((COMP_POINT--)); + done; + COMP_LINE=${COMP_LINE#"${COMP_WORDS[i]}"}; + ((COMP_POINT-=${#COMP_WORDS[i]})); + done; + for ((i=0; i <= COMP_CWORD - $word_offset; i++ )) + do + COMP_WORDS[i]=${COMP_WORDS[i+$word_offset]}; + done; + for ((i; i <= COMP_CWORD; i++ )) + do + unset 'COMP_WORDS[i]'; + done; + ((COMP_CWORD -= $word_offset)); + COMPREPLY=(); + local cur; + _get_comp_words_by_ref cur; + if [[ $COMP_CWORD -eq 0 ]]; then + local IFS=' +'; + compopt -o filenames; + COMPREPLY=($(compgen -d -c -- "$cur")); + else + local cmd=${COMP_WORDS[0]} compcmd=${COMP_WORDS[0]}; + local cspec=$(complete -p $cmd 2>/dev/null); + if [[ ! -n $cspec && $cmd == */* ]]; then + cspec=$(complete -p ${cmd##*/} 2>/dev/null); + [[ -n $cspec ]] && compcmd=${cmd##*/}; + fi; + if [[ ! -n $cspec ]]; then + compcmd=${cmd##*/}; + _completion_loader $compcmd; + cspec=$(complete -p $compcmd 2>/dev/null); + fi; + if [[ -n $cspec ]]; then + if [[ ${cspec#* -F } != $cspec ]]; then + local func=${cspec#*-F }; + func=${func%% *}; + if [[ ${#COMP_WORDS[@]} -ge 2 ]]; then + $func $cmd "${COMP_WORDS[${#COMP_WORDS[@]}-1]}" "${COMP_WORDS[${#COMP_WORDS[@]}-2]}"; + else + $func $cmd "${COMP_WORDS[${#COMP_WORDS[@]}-1]}"; + fi; + local opt; + while [[ $cspec == *" -o "* ]]; do + cspec=${cspec#*-o }; + opt=${cspec%% *}; + compopt -o $opt; + cspec=${cspec#$opt}; + done; + else + cspec=${cspec#complete}; + cspec=${cspec%%$compcmd}; + COMPREPLY=($(eval compgen "$cspec" -- '$cur')); + fi; + else + if [[ ${#COMPREPLY[@]} -eq 0 ]]; then + _minimal; + fi; + fi; + fi +} +_complete_as_root () +{ + [[ $EUID -eq 0 || -n ${root_command:-} ]] +} +_completion_loader () +{ + local cmd="${1:-_EmptycmD_}"; + __load_completion "$cmd" && return 124; + complete -F _minimal -- "$cmd" && return 124 +} +_configured_interfaces () +{ + if [[ -f /etc/debian_version ]]; then + COMPREPLY=($(compgen -W "$(command sed -ne 's|^iface \([^ ]\{1,\}\).*$|\1|p' /etc/network/interfaces /etc/network/interfaces.d/* 2>/dev/null)" -- "$cur")); + else + if [[ -f /etc/SuSE-release ]]; then + COMPREPLY=($(compgen -W "$(printf '%s\n' /etc/sysconfig/network/ifcfg-* | command sed -ne 's|.*ifcfg-\([^*].*\)$|\1|p')" -- "$cur")); + else + if [[ -f /etc/pld-release ]]; then + COMPREPLY=($(compgen -W "$(command ls -B /etc/sysconfig/interfaces | command sed -ne 's|.*ifcfg-\([^*].*\)$|\1|p')" -- "$cur")); + else + COMPREPLY=($(compgen -W "$(printf '%s\n' /etc/sysconfig/network-scripts/ifcfg-* | command sed -ne 's|.*ifcfg-\([^*].*\)$|\1|p')" -- "$cur")); + fi; + fi; + fi +} +_count_args () +{ + local i cword words; + __reassemble_comp_words_by_ref "$1" words cword; + args=1; + for ((i=1; i < cword; i++ )) + do + if [[ ${words[i]} != -* && ${words[i-1]} != $2 || ${words[i]} == $3 ]]; then + (( args++ )); + fi; + done +} +_dvd_devices () +{ + COMPREPLY+=($(compgen -f -d -X "!*/?(r)dvd*" -- "${cur:-/dev/}")) +} +_expand () +{ + if [[ "$cur" == \~*/* ]]; then + __expand_tilde_by_ref cur; + else + if [[ "$cur" == \~* ]]; then + _tilde "$cur" || eval COMPREPLY[0]=$(printf ~%q "${COMPREPLY[0]#\~}"); + return ${#COMPREPLY[@]}; + fi; + fi +} +_filedir () +{ + local IFS=' +'; + _tilde "$cur" || return; + local -a toks; + local reset; + if [[ "$1" == -d ]]; then + reset=$(shopt -po noglob); + set -o noglob; + toks=($(compgen -d -- "$cur")); + IFS=' '; + $reset; + IFS=' +'; + else + local quoted; + _quote_readline_by_ref "$cur" quoted; + local xspec=${1:+"!*.@($1|${1^^})"} plusdirs=(); + local opts=(-f -X "$xspec"); + [[ -n $xspec ]] && plusdirs=(-o plusdirs); + [[ -n ${COMP_FILEDIR_FALLBACK-} ]] || opts+=("${plusdirs[@]}"); + reset=$(shopt -po noglob); + set -o noglob; + toks+=($(compgen "${opts[@]}" -- $quoted)); + IFS=' '; + $reset; + IFS=' +'; + [[ -n ${COMP_FILEDIR_FALLBACK:-} && -n "$1" && ${#toks[@]} -lt 1 ]] && { + reset=$(shopt -po noglob); + set -o noglob; + toks+=($(compgen -f "${plusdirs[@]}" -- $quoted)); + IFS=' '; + $reset; + IFS=' +' + }; + fi; + if [[ ${#toks[@]} -ne 0 ]]; then + compopt -o filenames 2> /dev/null; + COMPREPLY+=("${toks[@]}"); + fi +} +_filedir_xspec () +{ + local cur prev words cword; + _init_completion || return; + _tilde "$cur" || return; + local IFS=' +' xspec=${_xspecs[${1##*/}]} tmp; + local -a toks; + toks=($( + compgen -d -- "$(quote_readline "$cur")" | { + while read -r tmp; do + printf '%s\n' $tmp + done + } + )); + eval xspec="${xspec}"; + local matchop=!; + if [[ $xspec == !* ]]; then + xspec=${xspec#!}; + matchop=@; + fi; + xspec="$matchop($xspec|${xspec^^})"; + toks+=($( + eval compgen -f -X "'!$xspec'" -- "\$(quote_readline "\$cur")" | { + while read -r tmp; do + [[ -n $tmp ]] && printf '%s\n' $tmp + done + } + )); + [[ -n ${COMP_FILEDIR_FALLBACK:-} && ${#toks[@]} -lt 1 ]] && { + local reset=$(shopt -po noglob); + set -o noglob; + toks+=($(compgen -f -- "$(quote_readline "$cur")")); + IFS=' '; + $reset; + IFS=' +' + }; + if [[ ${#toks[@]} -ne 0 ]]; then + compopt -o filenames; + COMPREPLY=("${toks[@]}"); + fi +} +_fstypes () +{ + local fss; + if [[ -e /proc/filesystems ]]; then + fss="$(cut -d' ' -f2 /proc/filesystems) + $(awk '! /\*/ { print $NF }' /etc/filesystems 2>/dev/null)"; + else + fss="$(awk '/^[ \t]*[^#]/ { print $3 }' /etc/fstab 2>/dev/null) + $(awk '/^[ \t]*[^#]/ { print $3 }' /etc/mnttab 2>/dev/null) + $(awk '/^[ \t]*[^#]/ { print $4 }' /etc/vfstab 2>/dev/null) + $(awk '{ print $1 }' /etc/dfs/fstypes 2>/dev/null) + $([[ -d /etc/fs ]] && command ls /etc/fs)"; + fi; + [[ -n $fss ]] && COMPREPLY+=($(compgen -W "$fss" -- "$cur")) +} +_get_comp_words_by_ref () +{ + local exclude flag i OPTIND=1; + local cur cword words=(); + local upargs=() upvars=() vcur vcword vprev vwords; + while getopts "c:i:n:p:w:" flag "$@"; do + case $flag in + c) + vcur=$OPTARG + ;; + i) + vcword=$OPTARG + ;; + n) + exclude=$OPTARG + ;; + p) + vprev=$OPTARG + ;; + w) + vwords=$OPTARG + ;; + esac; + done; + while [[ $# -ge $OPTIND ]]; do + case ${!OPTIND} in + cur) + vcur=cur + ;; + prev) + vprev=prev + ;; + cword) + vcword=cword + ;; + words) + vwords=words + ;; + *) + echo "bash_completion: $FUNCNAME: \`${!OPTIND}':" "unknown argument" 1>&2; + return 1 + ;; + esac; + (( OPTIND += 1 )); + done; + __get_cword_at_cursor_by_ref "$exclude" words cword cur; + [[ -n $vcur ]] && { + upvars+=("$vcur"); + upargs+=(-v $vcur "$cur") + }; + [[ -n $vcword ]] && { + upvars+=("$vcword"); + upargs+=(-v $vcword "$cword") + }; + [[ -n $vprev && $cword -ge 1 ]] && { + upvars+=("$vprev"); + upargs+=(-v $vprev "${words[cword - 1]}") + }; + [[ -n $vwords ]] && { + upvars+=("$vwords"); + upargs+=(-a${#words[@]} $vwords "${words[@]}") + }; + (( ${#upvars[@]} )) && local "${upvars[@]}" && _upvars "${upargs[@]}" +} +_get_cword () +{ + local LC_CTYPE=C; + local cword words; + __reassemble_comp_words_by_ref "$1" words cword; + if [[ -n ${2//[^0-9]/} ]]; then + printf "%s" "${words[cword-$2]}"; + else + if [[ "${#words[cword]}" -eq 0 || "$COMP_POINT" == "${#COMP_LINE}" ]]; then + printf "%s" "${words[cword]}"; + else + local i; + local cur="$COMP_LINE"; + local index="$COMP_POINT"; + for ((i = 0; i <= cword; ++i )) + do + while [[ "${#cur}" -ge ${#words[i]} && "${cur:0:${#words[i]}}" != "${words[i]}" ]]; do + cur="${cur:1}"; + [[ $index -gt 0 ]] && ((index--)); + done; + if [[ "$i" -lt "$cword" ]]; then + local old_size="${#cur}"; + cur="${cur#${words[i]}}"; + local new_size="${#cur}"; + (( index -= old_size - new_size )); + fi; + done; + if [[ "${words[cword]:0:${#cur}}" != "$cur" ]]; then + printf "%s" "${words[cword]}"; + else + printf "%s" "${cur:0:$index}"; + fi; + fi; + fi +} +_get_first_arg () +{ + local i; + arg=; + for ((i=1; i < COMP_CWORD; i++ )) + do + if [[ "${COMP_WORDS[i]}" != -* ]]; then + arg=${COMP_WORDS[i]}; + break; + fi; + done +} +_get_pword () +{ + if [[ $COMP_CWORD -ge 1 ]]; then + _get_cword "${@:-}" 1; + fi +} +_gids () +{ + if type getent &> /dev/null; then + COMPREPLY=($(compgen -W '$(getent group | cut -d: -f3)' -- "$cur")); + else + if type perl &> /dev/null; then + COMPREPLY=($(compgen -W '$(perl -e '"'"'while (($gid) = (getgrent)[2]) { print $gid . "\n" }'"'"')' -- "$cur")); + else + COMPREPLY=($(compgen -W '$(cut -d: -f3 /etc/group)' -- "$cur")); + fi; + fi +} +_have () +{ + PATH=$PATH:/usr/sbin:/sbin:/usr/local/sbin type $1 &> /dev/null +} +_included_ssh_config_files () +{ + [[ $# -lt 1 ]] && echo "bash_completion: $FUNCNAME: missing mandatory argument CONFIG" 1>&2; + local configfile i f; + configfile=$1; + local included=($(command sed -ne 's/^[[:blank:]]*[Ii][Nn][Cc][Ll][Uu][Dd][Ee][[:blank:]]\{1,\}\([^#%]*\)\(#.*\)\{0,1\}$/\1/p' "${configfile}")); + for i in "${included[@]}"; + do + if ! [[ "$i" =~ ^\~.*|^\/.* ]]; then + if [[ "$configfile" =~ ^\/etc\/ssh.* ]]; then + i="/etc/ssh/$i"; + else + i="$HOME/.ssh/$i"; + fi; + fi; + __expand_tilde_by_ref i; + for f in ${i}; + do + if [ -r $f ]; then + config+=("$f"); + _included_ssh_config_files $f; + fi; + done; + done +} +_init_completion () +{ + local exclude="" flag outx errx inx OPTIND=1; + while getopts "n:e:o:i:s" flag "$@"; do + case $flag in + n) + exclude+=$OPTARG + ;; + e) + errx=$OPTARG + ;; + o) + outx=$OPTARG + ;; + i) + inx=$OPTARG + ;; + s) + split=false; + exclude+== + ;; + esac; + done; + COMPREPLY=(); + local redir="@(?([0-9])<|?([0-9&])>?(>)|>&)"; + _get_comp_words_by_ref -n "$exclude<>&" cur prev words cword; + _variables && return 1; + if [[ $cur == $redir* || $prev == $redir ]]; then + local xspec; + case $cur in + 2'>'*) + xspec=$errx + ;; + *'>'*) + xspec=$outx + ;; + *'<'*) + xspec=$inx + ;; + *) + case $prev in + 2'>'*) + xspec=$errx + ;; + *'>'*) + xspec=$outx + ;; + *'<'*) + xspec=$inx + ;; + esac + ;; + esac; + cur="${cur##$redir}"; + _filedir $xspec; + return 1; + fi; + local i skip; + for ((i=1; i < ${#words[@]}; 1)) + do + if [[ ${words[i]} == $redir* ]]; then + [[ ${words[i]} == $redir ]] && skip=2 || skip=1; + words=("${words[@]:0:i}" "${words[@]:i+skip}"); + [[ $i -le $cword ]] && (( cword -= skip )); + else + (( i++ )); + fi; + done; + [[ $cword -le 0 ]] && return 1; + prev=${words[cword-1]}; + [[ -n ${split-} ]] && _split_longopt && split=true; + return 0 +} +_installed_modules () +{ + COMPREPLY=($(compgen -W "$(PATH="$PATH:/sbin" lsmod | awk '{if (NR != 1) print $1}')" -- "$1")) +} +_ip_addresses () +{ + local n; + case $1 in + -a) + n='6\?' + ;; + -6) + n='6' + ;; + esac; + local PATH=$PATH:/sbin; + local addrs=$({ LC_ALL=C ifconfig -a || ip addr show; } 2>/dev/null | + command sed -e 's/[[:space:]]addr:/ /' -ne "s|.*inet${n}[[:space:]]\{1,\}\([^[:space:]/]*\).*|\1|p"); + COMPREPLY+=($(compgen -W "$addrs" -- "$cur")) +} +_kernel_versions () +{ + COMPREPLY=($(compgen -W '$(command ls /lib/modules)' -- "$cur")) +} +_known_hosts () +{ + local cur prev words cword; + _init_completion -n : || return; + local options; + [[ "$1" == -a || "$2" == -a ]] && options=-a; + [[ "$1" == -c || "$2" == -c ]] && options+=" -c"; + _known_hosts_real $options -- "$cur" +} +_known_hosts_real () +{ + local configfile flag prefix OIFS=$IFS; + local cur user suffix aliases i host ipv4 ipv6; + local -a kh tmpkh khd config; + local OPTIND=1; + while getopts "ac46F:p:" flag "$@"; do + case $flag in + a) + aliases='yes' + ;; + c) + suffix=':' + ;; + F) + configfile=$OPTARG + ;; + p) + prefix=$OPTARG + ;; + 4) + ipv4=1 + ;; + 6) + ipv6=1 + ;; + esac; + done; + [[ $# -lt $OPTIND ]] && echo "bash_completion: $FUNCNAME: missing mandatory argument CWORD" 1>&2; + cur=${!OPTIND}; + (( OPTIND += 1 )); + [[ $# -ge $OPTIND ]] && echo "bash_completion: $FUNCNAME($*): unprocessed arguments:" $(while [[ $# -ge $OPTIND ]]; do printf '%s\n' ${!OPTIND}; shift; done) 1>&2; + [[ $cur == *@* ]] && user=${cur%@*}@ && cur=${cur#*@}; + kh=(); + if [[ -n $configfile ]]; then + [[ -r $configfile ]] && config+=("$configfile"); + else + for i in /etc/ssh/ssh_config ~/.ssh/config ~/.ssh2/config; + do + [[ -r $i ]] && config+=("$i"); + done; + fi; + for i in "${config[@]}"; + do + _included_ssh_config_files "$i"; + done; + if [[ ${#config[@]} -gt 0 ]]; then + local IFS=' +' j; + tmpkh=($(awk 'sub("^[ \t]*([Gg][Ll][Oo][Bb][Aa][Ll]|[Uu][Ss][Ee][Rr])[Kk][Nn][Oo][Ww][Nn][Hh][Oo][Ss][Tt][Ss][Ff][Ii][Ll][Ee][ \t]+", "") { print $0 }' "${config[@]}" | sort -u)); + IFS=$OIFS; + for i in "${tmpkh[@]}"; + do + while [[ $i =~ ^([^\"]*)\"([^\"]*)\"(.*)$ ]]; do + i=${BASH_REMATCH[1]}${BASH_REMATCH[3]}; + j=${BASH_REMATCH[2]}; + __expand_tilde_by_ref j; + [[ -r $j ]] && kh+=("$j"); + done; + for j in $i; + do + __expand_tilde_by_ref j; + [[ -r $j ]] && kh+=("$j"); + done; + done; + fi; + if [[ -z $configfile ]]; then + for i in /etc/ssh/ssh_known_hosts /etc/ssh/ssh_known_hosts2 /etc/known_hosts /etc/known_hosts2 ~/.ssh/known_hosts ~/.ssh/known_hosts2; + do + [[ -r $i ]] && kh+=("$i"); + done; + for i in /etc/ssh2/knownhosts ~/.ssh2/hostkeys; + do + [[ -d $i ]] && khd+=("$i"/*pub); + done; + fi; + if [[ ${#kh[@]} -gt 0 || ${#khd[@]} -gt 0 ]]; then + if [[ ${#kh[@]} -gt 0 ]]; then + for i in "${kh[@]}"; + do + while read -ra tmpkh; do + set -- "${tmpkh[@]}"; + [[ $1 == [\|\#]* ]] && continue; + [[ $1 == @* ]] && shift; + local IFS=,; + for host in $1; + do + [[ $host == *[*?]* ]] && continue; + host="${host#[}"; + host="${host%]?(:+([0-9]))}"; + COMPREPLY+=($host); + done; + IFS=$OIFS; + done < "$i"; + done; + COMPREPLY=($(compgen -W '${COMPREPLY[@]}' -- "$cur")); + fi; + if [[ ${#khd[@]} -gt 0 ]]; then + for i in "${khd[@]}"; + do + if [[ "$i" == *key_22_$cur*.pub && -r "$i" ]]; then + host=${i/#*key_22_/}; + host=${host/%.pub/}; + COMPREPLY+=($host); + fi; + done; + fi; + for ((i=0; i < ${#COMPREPLY[@]}; i++ )) + do + COMPREPLY[i]=$prefix$user${COMPREPLY[i]}$suffix; + done; + fi; + if [[ ${#config[@]} -gt 0 && -n "$aliases" ]]; then + local hosts=$(command sed -ne 's/^[[:blank:]]*[Hh][Oo][Ss][Tt][[:blank:]]\{1,\}\([^#*?%]*\)\(#.*\)\{0,1\}$/\1/p' "${config[@]}"); + COMPREPLY+=($(compgen -P "$prefix$user" -S "$suffix" -W "$hosts" -- "$cur")); + fi; + if [[ -n ${COMP_KNOWN_HOSTS_WITH_AVAHI:-} ]] && type avahi-browse &> /dev/null; then + COMPREPLY+=($(compgen -P "$prefix$user" -S "$suffix" -W "$(avahi-browse -cpr _workstation._tcp 2>/dev/null | awk -F';' '/^=/ { print $7 }' | sort -u)" -- "$cur")); + fi; + COMPREPLY+=($(compgen -W "$(ruptime 2>/dev/null | awk '!/^ruptime:/ { print $1 }')" -- "$cur")); + if [[ -n ${COMP_KNOWN_HOSTS_WITH_HOSTFILE-1} ]]; then + COMPREPLY+=($(compgen -A hostname -P "$prefix$user" -S "$suffix" -- "$cur")); + fi; + if [[ -n $ipv4 ]]; then + COMPREPLY=("${COMPREPLY[@]/*:*$suffix/}"); + fi; + if [[ -n $ipv6 ]]; then + COMPREPLY=("${COMPREPLY[@]/+([0-9]).+([0-9]).+([0-9]).+([0-9])$suffix/}"); + fi; + if [[ -n $ipv4 || -n $ipv6 ]]; then + for i in "${!COMPREPLY[@]}"; + do + [[ -n ${COMPREPLY[i]} ]] || unset -v COMPREPLY[i]; + done; + fi; + __ltrim_colon_completions "$prefix$user$cur" +} +_longopt () +{ + local cur prev words cword split; + _init_completion -s || return; + case "${prev,,}" in + --help | --usage | --version) + return + ;; + --!(no-*)dir*) + _filedir -d; + return + ;; + --!(no-*)@(file|path)*) + _filedir; + return + ;; + --+([-a-z0-9_])) + local argtype=$(LC_ALL=C $1 --help 2>&1 | command sed -ne "s|.*$prev\[\{0,1\}=[<[]\{0,1\}\([-A-Za-z0-9_]\{1,\}\).*|\1|p"); + case ${argtype,,} in + *dir*) + _filedir -d; + return + ;; + *file* | *path*) + _filedir; + return + ;; + esac + ;; + esac; + $split && return; + if [[ "$cur" == -* ]]; then + COMPREPLY=($(compgen -W "$(LC_ALL=C $1 --help 2>&1 | while read -r line; do [[ $line =~ --[-A-Za-z0-9]+=? ]] && printf '%s\n' ${BASH_REMATCH[0]} + done)" -- "$cur")); + [[ $COMPREPLY == *= ]] && compopt -o nospace; + else + if [[ "$1" == *@(rmdir|chroot) ]]; then + _filedir -d; + else + [[ "$1" == *mkdir ]] && compopt -o nospace; + _filedir; + fi; + fi +} +_mac_addresses () +{ + local re='\([A-Fa-f0-9]\{2\}:\)\{5\}[A-Fa-f0-9]\{2\}'; + local PATH="$PATH:/sbin:/usr/sbin"; + COMPREPLY+=($( { LC_ALL=C ifconfig -a || ip link show; } 2>/dev/null | command sed -ne "s/.*[[:space:]]HWaddr[[:space:]]\{1,\}\($re\)[[:space:]].*/\1/p" -ne "s/.*[[:space:]]HWaddr[[:space:]]\{1,\}\($re\)[[:space:]]*$/\1/p" -ne "s|.*[[:space:]]\(link/\)\{0,1\}ether[[:space:]]\{1,\}\($re\)[[:space:]].*|\2|p" -ne "s|.*[[:space:]]\(link/\)\{0,1\}ether[[:space:]]\{1,\}\($re\)[[:space:]]*$|\2|p" + )); + COMPREPLY+=($({ arp -an || ip neigh show; } 2>/dev/null | command sed -ne "s/.*[[:space:]]\($re\)[[:space:]].*/\1/p" -ne "s/.*[[:space:]]\($re\)[[:space:]]*$/\1/p")); + COMPREPLY+=($(command sed -ne "s/^[[:space:]]*\($re\)[[:space:]].*/\1/p" /etc/ethers 2>/dev/null)); + COMPREPLY=($(compgen -W '${COMPREPLY[@]}' -- "$cur")); + __ltrim_colon_completions "$cur" +} +_minimal () +{ + local cur prev words cword split; + _init_completion -s || return; + $split && return; + _filedir +} +_modules () +{ + local modpath; + modpath=/lib/modules/$1; + COMPREPLY=($(compgen -W "$(command ls -RL $modpath 2>/dev/null | command sed -ne 's/^\(.*\)\.k\{0,1\}o\(\.[gx]z\)\{0,1\}$/\1/p')" -- "$cur")) +} +_ncpus () +{ + local var=NPROCESSORS_ONLN; + [[ $OSTYPE == *linux* ]] && var=_$var; + local n=$(getconf $var 2>/dev/null); + printf %s ${n:-1} +} +_parse_help () +{ + eval local cmd=$(quote "$1"); + local line; + { + case $cmd in + -) + cat + ;; + *) + LC_ALL=C "$(dequote "$cmd")" ${2:---help} 2>&1 + ;; + esac + } | while read -r line; do + [[ $line == *([[:blank:]])-* ]] || continue; + while [[ $line =~ ((^|[^-])-[A-Za-z0-9?][[:space:]]+)\[?[A-Z0-9]+([,_-]+[A-Z0-9]+)?(\.\.+)?\]? ]]; do + line=${line/"${BASH_REMATCH[0]}"/"${BASH_REMATCH[1]}"}; + done; + __parse_options "${line// or /, }"; + done +} +_parse_usage () +{ + eval local cmd=$(quote "$1"); + local line match option i char; + { + case $cmd in + -) + cat + ;; + *) + LC_ALL=C "$(dequote "$cmd")" ${2:---usage} 2>&1 + ;; + esac + } | while read -r line; do + while [[ $line =~ \[[[:space:]]*(-[^]]+)[[:space:]]*\] ]]; do + match=${BASH_REMATCH[0]}; + option=${BASH_REMATCH[1]}; + case $option in + -?(\[)+([a-zA-Z0-9?])) + for ((i=1; i < ${#option}; i++ )) + do + char=${option:i:1}; + [[ $char != '[' ]] && printf '%s\n' -$char; + done + ;; + *) + __parse_options "$option" + ;; + esac; + line=${line#*"$match"}; + done; + done +} +_pci_ids () +{ + COMPREPLY+=($(compgen -W "$(PATH="$PATH:/sbin" lspci -n | awk '{print $3}')" -- "$cur")) +} +_pgids () +{ + COMPREPLY=($(compgen -W '$(command ps axo pgid=)' -- "$cur")) +} +_pids () +{ + COMPREPLY=($(compgen -W '$(command ps axo pid=)' -- "$cur")) +} +_pnames () +{ + local -a procs; + if [[ "$1" == -s ]]; then + procs=($(command ps axo comm | command sed -e 1d)); + else + local line i=-1 OIFS=$IFS; + IFS=' +'; + local -a psout=($(command ps axo command=)); + IFS=$OIFS; + for line in "${psout[@]}"; + do + if [[ $i -eq -1 ]]; then + if [[ $line =~ ^(.*[[:space:]])COMMAND([[:space:]]|$) ]]; then + i=${#BASH_REMATCH[1]}; + else + break; + fi; + else + line=${line:$i}; + line=${line%% *}; + procs+=($line); + fi; + done; + if [[ $i -eq -1 ]]; then + for line in "${psout[@]}"; + do + if [[ $line =~ ^[[(](.+)[])]$ ]]; then + procs+=(${BASH_REMATCH[1]}); + else + line=${line%% *}; + line=${line##@(*/|-)}; + procs+=($line); + fi; + done; + fi; + fi; + COMPREPLY=($(compgen -X "" -W '${procs[@]}' -- "$cur" )) +} +_quote_readline_by_ref () +{ + if [ -z "$1" ]; then + printf -v $2 %s "$1"; + else + if [[ $1 == \'* ]]; then + printf -v $2 %s "${1:1}"; + else + if [[ $1 == \~* ]]; then + printf -v $2 \~%q "${1:1}"; + else + printf -v $2 %q "$1"; + fi; + fi; + fi; + [[ ${!2} == \$* ]] && eval $2=${!2} +} +_realcommand () +{ + type -P "$1" > /dev/null && { + if type -p realpath > /dev/null; then + realpath "$(type -P "$1")"; + else + if type -p greadlink > /dev/null; then + greadlink -f "$(type -P "$1")"; + else + if type -p readlink > /dev/null; then + readlink -f "$(type -P "$1")"; + else + type -P "$1"; + fi; + fi; + fi + } +} +_rl_enabled () +{ + [[ "$(bind -v)" == *$1+([[:space:]])on* ]] +} +_root_command () +{ + local PATH=$PATH:/sbin:/usr/sbin:/usr/local/sbin; + local root_command=$1; + _command +} +_service () +{ + local cur prev words cword; + _init_completion || return; + [[ $cword -gt 2 ]] && return; + if [[ $cword -eq 1 && $prev == ?(*/)service ]]; then + _services; + [[ -e /etc/mandrake-release ]] && _xinetd_services; + else + local sysvdirs; + _sysvdirs; + COMPREPLY=($(compgen -W '`command sed -e "y/|/ /" \ + -ne "s/^.*\(U\|msg_u\)sage.*{\(.*\)}.*$/\2/p" \ + ${sysvdirs[0]}/${prev##*/} 2>/dev/null` start stop' -- "$cur")); + fi +} +_services () +{ + local sysvdirs; + _sysvdirs; + local IFS=' +' reset=$(shopt -p nullglob); + shopt -s nullglob; + COMPREPLY=($(printf '%s\n' ${sysvdirs[0]}/!($_backup_glob|functions|README))); + $reset; + COMPREPLY+=($({ systemctl list-units --full --all || systemctl list-unit-files; } 2>/dev/null | awk '$1 ~ /\.service$/ { sub("\\.service$", "", $1); print $1 }')); + if [[ -x /sbin/upstart-udev-bridge ]]; then + COMPREPLY+=($(initctl list 2>/dev/null | cut -d' ' -f1)); + fi; + COMPREPLY=($(compgen -W '${COMPREPLY[@]#${sysvdirs[0]}/}' -- "$cur")) +} +_shells () +{ + local shell rest; + while read -r shell rest; do + [[ $shell == /* && $shell == "$cur"* ]] && COMPREPLY+=($shell); + done 2> /dev/null < /etc/shells +} +_signals () +{ + local -a sigs=($(compgen -P "$1" -A signal "SIG${cur#$1}")); + COMPREPLY+=("${sigs[@]/#${1}SIG/${1}}") +} +_split_longopt () +{ + if [[ "$cur" == --?*=* ]]; then + prev="${cur%%?(\\)=*}"; + cur="${cur#*=}"; + return 0; + fi; + return 1 +} +_sysvdirs () +{ + sysvdirs=(); + [[ -d /etc/rc.d/init.d ]] && sysvdirs+=(/etc/rc.d/init.d); + [[ -d /etc/init.d ]] && sysvdirs+=(/etc/init.d); + [[ -f /etc/slackware-version ]] && sysvdirs=(/etc/rc.d); + return 0 +} +_terms () +{ + COMPREPLY+=($(compgen -W "$({ command sed -ne 's/^\([^[:space:]#|]\{2,\}\)|.*/\1/p' /etc/termcap; + { toe -a || toe; } | awk '{ print $1 }'; + find /{etc,lib,usr/lib,usr/share}/terminfo/? -type f -maxdepth 1 | awk -F/ '{ print $NF }'; + } 2>/dev/null)" -- "$cur")) +} +_tilde () +{ + local result=0; + if [[ $1 == \~* && $1 != */* ]]; then + COMPREPLY=($(compgen -P '~' -u -- "${1#\~}")); + result=${#COMPREPLY[@]}; + [[ $result -gt 0 ]] && compopt -o filenames 2> /dev/null; + fi; + return $result +} +_uids () +{ + if type getent &> /dev/null; then + COMPREPLY=($(compgen -W '$(getent passwd | cut -d: -f3)' -- "$cur")); + else + if type perl &> /dev/null; then + COMPREPLY=($(compgen -W '$(perl -e '"'"'while (($uid) = (getpwent)[2]) { print $uid . "\n" }'"'"')' -- "$cur")); + else + COMPREPLY=($(compgen -W '$(cut -d: -f3 /etc/passwd)' -- "$cur")); + fi; + fi +} +_upvar () +{ + echo "bash_completion: $FUNCNAME: deprecated function," "use _upvars instead" 1>&2; + if unset -v "$1"; then + if (( $# == 2 )); then + eval $1=\"\$2\"; + else + eval $1=\(\"\${@:2}\"\); + fi; + fi +} +_upvars () +{ + if ! (( $# )); then + echo "bash_completion: $FUNCNAME: usage: $FUNCNAME" "[-v varname value] | [-aN varname [value ...]] ..." 1>&2; + return 2; + fi; + while (( $# )); do + case $1 in + -a*) + [[ -n ${1#-a} ]] || { + echo "bash_completion: $FUNCNAME:" "\`$1': missing number specifier" 1>&2; + return 1 + }; + printf %d "${1#-a}" &> /dev/null || { + echo bash_completion: "$FUNCNAME: \`$1': invalid number specifier" 1>&2; + return 1 + }; + [[ -n "$2" ]] && unset -v "$2" && eval $2=\(\"\${@:3:${1#-a}}\"\) && shift $((${1#-a} + 2)) || { + echo bash_completion: "$FUNCNAME: \`$1${2+ }$2': missing argument(s)" 1>&2; + return 1 + } + ;; + -v) + [[ -n "$2" ]] && unset -v "$2" && eval $2=\"\$3\" && shift 3 || { + echo "bash_completion: $FUNCNAME: $1:" "missing argument(s)" 1>&2; + return 1 + } + ;; + *) + echo "bash_completion: $FUNCNAME: $1: invalid option" 1>&2; + return 1 + ;; + esac; + done +} +_usb_ids () +{ + COMPREPLY+=($(compgen -W "$(PATH="$PATH:/sbin" lsusb | awk '{print $6}')" -- "$cur")) +} +_user_at_host () +{ + local cur prev words cword; + _init_completion -n : || return; + if [[ $cur == *@* ]]; then + _known_hosts_real "$cur"; + else + COMPREPLY=($(compgen -u -S @ -- "$cur")); + compopt -o nospace; + fi +} +_usergroup () +{ + if [[ $cur == *\\\\* || $cur == *:*:* ]]; then + return; + else + if [[ $cur == *\\:* ]]; then + local prefix; + prefix=${cur%%*([^:])}; + prefix=${prefix//\\}; + local mycur="${cur#*[:]}"; + if [[ $1 == -u ]]; then + _allowed_groups "$mycur"; + else + local IFS=' +'; + COMPREPLY=($(compgen -g -- "$mycur")); + fi; + COMPREPLY=($(compgen -P "$prefix" -W "${COMPREPLY[@]}")); + else + if [[ $cur == *:* ]]; then + local mycur="${cur#*:}"; + if [[ $1 == -u ]]; then + _allowed_groups "$mycur"; + else + local IFS=' +'; + COMPREPLY=($(compgen -g -- "$mycur")); + fi; + else + if [[ $1 == -u ]]; then + _allowed_users "$cur"; + else + local IFS=' +'; + COMPREPLY=($(compgen -u -- "$cur")); + fi; + fi; + fi; + fi +} +_userland () +{ + local userland=$(uname -s); + [[ $userland == @(Linux|GNU/*) ]] && userland=GNU; + [[ $userland == $1 ]] +} +_variables () +{ + if [[ $cur =~ ^(\$(\{[!#]?)?)([A-Za-z0-9_]*)$ ]]; then + if [[ $cur == \${* ]]; then + local arrs vars; + vars=($(compgen -A variable -P ${BASH_REMATCH[1]} -S '}' -- ${BASH_REMATCH[3]})) && arrs=($(compgen -A arrayvar -P ${BASH_REMATCH[1]} -S '[' -- ${BASH_REMATCH[3]})); + if [[ ${#vars[@]} -eq 1 && -n $arrs ]]; then + compopt -o nospace; + COMPREPLY+=(${arrs[*]}); + else + COMPREPLY+=(${vars[*]}); + fi; + else + COMPREPLY+=($(compgen -A variable -P '$' -- "${BASH_REMATCH[3]}")); + fi; + return 0; + else + if [[ $cur =~ ^(\$\{[#!]?)([A-Za-z0-9_]*)\[([^]]*)$ ]]; then + local IFS=' +'; + COMPREPLY+=($(compgen -W '$(printf %s\\n "${!'${BASH_REMATCH[2]}'[@]}")' -P "${BASH_REMATCH[1]}${BASH_REMATCH[2]}[" -S ']}' -- "${BASH_REMATCH[3]}")); + if [[ ${BASH_REMATCH[3]} == [@*] ]]; then + COMPREPLY+=("${BASH_REMATCH[1]}${BASH_REMATCH[2]}[${BASH_REMATCH[3]}]}"); + fi; + __ltrim_colon_completions "$cur"; + return 0; + else + if [[ $cur =~ ^\$\{[#!]?[A-Za-z0-9_]*\[.*\]$ ]]; then + COMPREPLY+=("$cur}"); + __ltrim_colon_completions "$cur"; + return 0; + else + case $prev in + TZ) + cur=/usr/share/zoneinfo/$cur; + _filedir; + for i in "${!COMPREPLY[@]}"; + do + if [[ ${COMPREPLY[i]} == *.tab ]]; then + unset 'COMPREPLY[i]'; + continue; + else + if [[ -d ${COMPREPLY[i]} ]]; then + COMPREPLY[i]+=/; + compopt -o nospace; + fi; + fi; + COMPREPLY[i]=${COMPREPLY[i]#/usr/share/zoneinfo/}; + done; + return 0 + ;; + TERM) + _terms; + return 0 + ;; + LANG | LC_*) + COMPREPLY=($(compgen -W '$(locale -a 2>/dev/null)' -- "$cur" )); + return 0 + ;; + esac; + fi; + fi; + fi; + return 1 +} +_xfunc () +{ + set -- "$@"; + local srcfile=$1; + shift; + declare -F $1 &> /dev/null || { + __load_completion "$srcfile" + }; + "$@" +} +_xinetd_services () +{ + local xinetddir=/etc/xinetd.d; + if [[ -d $xinetddir ]]; then + local IFS=' +' reset=$(shopt -p nullglob); + shopt -s nullglob; + local -a svcs=($(printf '%s\n' $xinetddir/!($_backup_glob))); + $reset; + COMPREPLY+=($(compgen -W '${svcs[@]#$xinetddir/}' -- "$cur")); + fi +} +command_not_found_handle () +{ + if [ -x /usr/lib/command-not-found ]; then + /usr/lib/command-not-found -- "$1"; + return $?; + else + if [ -x /usr/share/command-not-found/command-not-found ]; then + /usr/share/command-not-found/command-not-found -- "$1"; + return $?; + else + printf "%s: command not found\n" "$1" 1>&2; + return 127; + fi; + fi +} +dequote () +{ + eval printf %s "$1" 2> /dev/null +} +gawklibpath_append () +{ + [ -z "$AWKLIBPATH" ] && AWKLIBPATH=`gawk 'BEGIN {print ENVIRON["AWKLIBPATH"]}'`; + export AWKLIBPATH="$AWKLIBPATH:$*" +} +gawklibpath_default () +{ + unset AWKLIBPATH; + export AWKLIBPATH=`gawk 'BEGIN {print ENVIRON["AWKLIBPATH"]}'` +} +gawklibpath_prepend () +{ + [ -z "$AWKLIBPATH" ] && AWKLIBPATH=`gawk 'BEGIN {print ENVIRON["AWKLIBPATH"]}'`; + export AWKLIBPATH="$*:$AWKLIBPATH" +} +gawkpath_append () +{ + [ -z "$AWKPATH" ] && AWKPATH=`gawk 'BEGIN {print ENVIRON["AWKPATH"]}'`; + export AWKPATH="$AWKPATH:$*" +} +gawkpath_default () +{ + unset AWKPATH; + export AWKPATH=`gawk 'BEGIN {print ENVIRON["AWKPATH"]}'` +} +gawkpath_prepend () +{ + [ -z "$AWKPATH" ] && AWKPATH=`gawk 'BEGIN {print ENVIRON["AWKPATH"]}'`; + export AWKPATH="$*:$AWKPATH" +} +quote () +{ + local quoted=${1//\'/\'\\\'\'}; + printf "'%s'" "$quoted" +} +quote_readline () +{ + local quoted; + _quote_readline_by_ref "$1" ret; + printf %s "$ret" +} diff --git a/wavesrv/pkg/cmdrunner/resolver.go b/wavesrv/pkg/cmdrunner/resolver.go new file mode 100644 index 00000000..7142a369 --- /dev/null +++ b/wavesrv/pkg/cmdrunner/resolver.go @@ -0,0 +1,516 @@ +package cmdrunner + +import ( + "context" + "fmt" + "log" + "regexp" + "strconv" + "strings" + + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/scpacket" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/google/uuid" +) + +const ( + R_Session = 1 + R_Screen = 2 + R_Remote = 8 + R_RemoteConnected = 16 +) + +type resolvedIds struct { + SessionId string + ScreenId string + Remote *ResolvedRemote +} + +type ResolvedRemote struct { + DisplayName string + RemotePtr sstore.RemotePtrType + MShell *remote.MShellProc + RState remote.RemoteRuntimeState + RemoteCopy *sstore.RemoteType + StatePtr *sstore.ShellStatePtr + FeState map[string]string +} + +type ResolveItem = sstore.ResolveItem + +func itemNames(items []ResolveItem) []string { + if len(items) == 0 { + return nil + } + rtn := make([]string, len(items)) + for idx, item := range items { + rtn[idx] = item.Name + } + return rtn +} + +func sessionsToResolveItems(sessions []*sstore.SessionType) []ResolveItem { + if len(sessions) == 0 { + return nil + } + rtn := make([]ResolveItem, len(sessions)) + for idx, session := range sessions { + rtn[idx] = ResolveItem{Name: session.Name, Id: session.SessionId, Hidden: session.Archived} + } + return rtn +} + +func screensToResolveItems(screens []*sstore.ScreenType) []ResolveItem { + if len(screens) == 0 { + return nil + } + rtn := make([]ResolveItem, len(screens)) + for idx, screen := range screens { + rtn[idx] = ResolveItem{Name: screen.Name, Id: screen.ScreenId, Hidden: screen.Archived} + } + return rtn +} + +// 1-indexed +func boundInt(ival int, maxVal int, wrap bool) int { + if maxVal == 0 { + return 0 + } + if ival < 1 { + if wrap { + return maxVal + } else { + return 1 + } + } + if ival > maxVal { + if wrap { + return 1 + } else { + return maxVal + } + } + return ival +} + +type posArgType struct { + Pos int + IsWrap bool + IsRelative bool + StartAnchor bool + EndAnchor bool +} + +func parsePosArg(posStr string) *posArgType { + if !positionRe.MatchString(posStr) { + return nil + } + if posStr == "+" { + return &posArgType{Pos: 1, IsWrap: true, IsRelative: true} + } else if posStr == "-" { + return &posArgType{Pos: -1, IsWrap: true, IsRelative: true} + } else if posStr == "S" { + return &posArgType{Pos: 0, IsRelative: true, StartAnchor: true} + } else if posStr == "E" { + return &posArgType{Pos: 0, IsRelative: true, EndAnchor: true} + } + if strings.HasPrefix(posStr, "S+") { + pos, _ := strconv.Atoi(posStr[2:]) + return &posArgType{Pos: pos, IsRelative: true, StartAnchor: true} + } + if strings.HasPrefix(posStr, "E-") { + pos, _ := strconv.Atoi(posStr[1:]) + return &posArgType{Pos: pos, IsRelative: true, EndAnchor: true} + } + if strings.HasPrefix(posStr, "+") || strings.HasPrefix(posStr, "-") { + pos, _ := strconv.Atoi(posStr) + return &posArgType{Pos: pos, IsRelative: true} + } + pos, _ := strconv.Atoi(posStr) + return &posArgType{Pos: pos} +} + +func resolveByPosition(isNumeric bool, allItems []ResolveItem, curId string, posStr string) *ResolveItem { + items := make([]ResolveItem, 0, len(allItems)) + for _, item := range allItems { + if !item.Hidden { + items = append(items, item) + } + } + if len(items) == 0 { + return nil + } + posArg := parsePosArg(posStr) + if posArg == nil { + return nil + } + var finalPos int + if posArg.IsRelative { + var curIdx int + if posArg.StartAnchor { + curIdx = 1 + } else if posArg.EndAnchor { + curIdx = len(items) + } else { + curIdx = 1 // if no match, curIdx will be first item + for idx, item := range items { + if item.Id == curId { + curIdx = idx + 1 + break + } + } + } + finalPos = curIdx + posArg.Pos + finalPos = boundInt(finalPos, len(items), posArg.IsWrap) + return &items[finalPos-1] + } else if isNumeric { + // these resolve items have a "Num" set that should be used to look up non-relative positions + // use allItems for numeric resolve + for _, item := range allItems { + if item.Num == posArg.Pos { + return &item + } + } + return nil + } else { + // non-numeric means position is just the index + finalPos = posArg.Pos + if finalPos <= 0 || finalPos > len(items) { + return nil + } + return &items[finalPos-1] + } +} + +func resolveRemoteArg(remoteArg string) (*sstore.RemotePtrType, error) { + rrUser, rrRemote, rrName, err := parseFullRemoteRef(remoteArg) + if err != nil { + return nil, err + } + if rrUser != "" { + return nil, fmt.Errorf("remoteusers not supported") + } + msh := remote.GetRemoteByArg(rrRemote) + if msh == nil { + return nil, nil + } + rcopy := msh.GetRemoteCopy() + return &sstore.RemotePtrType{RemoteId: rcopy.RemoteId, Name: rrName}, nil +} + +func resolveUiIds(ctx context.Context, pk *scpacket.FeCommandPacketType, rtype int) (resolvedIds, error) { + rtn := resolvedIds{} + uictx := pk.UIContext + if uictx != nil { + rtn.SessionId = uictx.SessionId + rtn.ScreenId = uictx.ScreenId + } + if pk.Kwargs["session"] != "" { + sessionId, err := resolveSessionArg(pk.Kwargs["session"]) + if err != nil { + return rtn, err + } + if sessionId != "" { + rtn.SessionId = sessionId + } + } + if pk.Kwargs["screen"] != "" { + screenId, err := resolveScreenArg(rtn.SessionId, pk.Kwargs["screen"]) + if err != nil { + return rtn, err + } + if screenId != "" { + rtn.ScreenId = screenId + } + } + var rptr *sstore.RemotePtrType + var err error + if pk.Kwargs["remote"] != "" { + rptr, err = resolveRemoteArg(pk.Kwargs["remote"]) + if err != nil { + return rtn, err + } + if rptr == nil { + return rtn, fmt.Errorf("invalid remote argument %q passed, remote not found", pk.Kwargs["remote"]) + } + } else if uictx.Remote != nil { + rptr = uictx.Remote + } + if rptr != nil { + err = rptr.Validate() + if err != nil { + return rtn, fmt.Errorf("invalid resolved remote: %v", err) + } + rr, err := ResolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.ScreenId) + if err != nil { + return rtn, err + } + rtn.Remote = rr + } + if rtype&R_Session > 0 && rtn.SessionId == "" { + return rtn, fmt.Errorf("no session") + } + if rtype&R_Screen > 0 && rtn.ScreenId == "" { + return rtn, fmt.Errorf("no screen") + } + if (rtype&R_Remote > 0 || rtype&R_RemoteConnected > 0) && rtn.Remote == nil { + return rtn, fmt.Errorf("no remote") + } + if rtype&R_RemoteConnected > 0 { + if !rtn.Remote.RState.IsConnected() { + err = rtn.Remote.MShell.TryAutoConnect() + if err != nil { + return rtn, fmt.Errorf("error trying to auto-connect remote [%s]: %w", rtn.Remote.DisplayName, err) + } + rrNew, err := ResolveRemoteFromPtr(ctx, rptr, rtn.SessionId, rtn.ScreenId) + if err != nil { + return rtn, err + } + rtn.Remote = rrNew + } + if !rtn.Remote.RState.IsConnected() { + return rtn, fmt.Errorf("remote [%s] is not connected", rtn.Remote.DisplayName) + } + if rtn.Remote.StatePtr == nil || rtn.Remote.FeState == nil { + return rtn, fmt.Errorf("remote [%s] state is not available", rtn.Remote.DisplayName) + } + } + return rtn, nil +} + +func resolveSessionScreen(ctx context.Context, sessionId string, screenArg string, curScreenArg string) (*ResolveItem, error) { + screens, err := sstore.GetSessionScreens(ctx, sessionId) + if err != nil { + return nil, fmt.Errorf("could not retreive screens for session=%s: %v", sessionId, err) + } + ritems := screensToResolveItems(screens) + return genericResolve(screenArg, curScreenArg, ritems, false, "screen") +} + +func resolveSession(ctx context.Context, sessionArg string, curSessionArg string) (*ResolveItem, error) { + bareSessions, err := sstore.GetBareSessions(ctx) + if err != nil { + return nil, err + } + ritems := sessionsToResolveItems(bareSessions) + ritem, err := genericResolve(sessionArg, curSessionArg, ritems, false, "session") + if err != nil { + return nil, err + } + return ritem, nil +} + +func resolveLine(ctx context.Context, sessionId string, screenId string, lineArg string, curLineArg string) (*ResolveItem, error) { + lines, err := sstore.GetLineResolveItems(ctx, screenId) + if err != nil { + return nil, fmt.Errorf("could not get lines: %v", err) + } + return genericResolve(lineArg, curLineArg, lines, true, "line") +} + +func getSessionIds(sarr []*sstore.SessionType) []string { + rtn := make([]string, len(sarr)) + for idx, s := range sarr { + rtn[idx] = s.SessionId + } + return rtn +} + +var partialUUIDRe = regexp.MustCompile("^[0-9a-f]{8}$") + +func isPartialUUID(s string) bool { + return partialUUIDRe.MatchString(s) +} + +func isUUID(s string) bool { + _, err := uuid.Parse(s) + return err == nil +} + +func getResolveItemById(id string, items []ResolveItem) *ResolveItem { + if id == "" { + return nil + } + for _, item := range items { + if item.Id == id { + return &item + } + } + return nil +} + +func genericResolve(arg string, curArg string, items []ResolveItem, isNumeric bool, typeStr string) (*ResolveItem, error) { + if len(items) == 0 || arg == "" { + return nil, nil + } + var curId string + if curArg != "" { + curItem, _ := genericResolve(curArg, "", items, isNumeric, typeStr) + if curItem != nil { + curId = curItem.Id + } + } + rtnItem := resolveByPosition(isNumeric, items, curId, arg) + if rtnItem != nil { + return rtnItem, nil + } + isUuid := isUUID(arg) + tryPuid := isPartialUUID(arg) + var prefixMatches []ResolveItem + for _, item := range items { + if (isUuid && item.Id == arg) || (tryPuid && strings.HasPrefix(item.Id, arg)) { + return &item, nil + } + if item.Name != "" { + if item.Name == arg { + return &item, nil + } + if !item.Hidden && strings.HasPrefix(item.Name, arg) { + prefixMatches = append(prefixMatches, item) + } + } + } + if len(prefixMatches) == 1 { + return &prefixMatches[0], nil + } + if len(prefixMatches) > 1 { + return nil, fmt.Errorf("could not resolve %s '%s', ambiguious prefix matched multiple %ss: %s", typeStr, arg, typeStr, formatStrs(itemNames(prefixMatches), "and", true)) + } + return nil, fmt.Errorf("could not resolve %s '%s' (name/id/pos not found)", typeStr, arg) +} + +func resolveSessionId(pk *scpacket.FeCommandPacketType) (string, error) { + sessionId := pk.Kwargs["session"] + if sessionId == "" { + return "", nil + } + if _, err := uuid.Parse(sessionId); err != nil { + return "", fmt.Errorf("invalid sessionid '%s'", sessionId) + } + return sessionId, nil +} + +func resolveSessionArg(sessionArg string) (string, error) { + if sessionArg == "" { + return "", nil + } + if _, err := uuid.Parse(sessionArg); err != nil { + return "", fmt.Errorf("invalid session arg specified (must be sessionid) '%s'", sessionArg) + } + return sessionArg, nil +} + +func resolveScreenArg(sessionId string, screenArg string) (string, error) { + if screenArg == "" { + return "", nil + } + if _, err := uuid.Parse(screenArg); err != nil { + return "", fmt.Errorf("invalid screen arg specified (must be screenid) '%s'", screenArg) + } + return screenArg, nil +} + +func resolveScreenId(ctx context.Context, pk *scpacket.FeCommandPacketType, sessionId string) (string, error) { + screenArg := pk.Kwargs["screen"] + if screenArg == "" { + return "", nil + } + if _, err := uuid.Parse(screenArg); err == nil { + return screenArg, nil + } + if sessionId == "" { + return "", fmt.Errorf("cannot resolve screen without session") + } + ritem, err := resolveSessionScreen(ctx, sessionId, screenArg, "") + if err != nil { + return "", err + } + return ritem.Id, nil +} + +// returns (remoteuserref, remoteref, name, error) +func parseFullRemoteRef(fullRemoteRef string) (string, string, string, error) { + if strings.HasPrefix(fullRemoteRef, "[") && strings.HasSuffix(fullRemoteRef, "]") { + fullRemoteRef = fullRemoteRef[1 : len(fullRemoteRef)-1] + } + fields := strings.Split(fullRemoteRef, ":") + if len(fields) > 3 { + return "", "", "", fmt.Errorf("invalid remote format '%s'", fullRemoteRef) + } + if len(fields) == 1 { + return "", fields[0], "", nil + } + if len(fields) == 2 { + if strings.HasPrefix(fields[0], "@") { + return fields[0], fields[1], "", nil + } + return "", fields[0], fields[1], nil + } + return fields[0], fields[1], fields[2], nil +} + +func ResolveRemoteFromPtr(ctx context.Context, rptr *sstore.RemotePtrType, sessionId string, screenId string) (*ResolvedRemote, error) { + if rptr == nil || rptr.RemoteId == "" { + return nil, nil + } + msh := remote.GetRemoteById(rptr.RemoteId) + if msh == nil { + return nil, fmt.Errorf("invalid remote '%s', not found", rptr.RemoteId) + } + rstate := msh.GetRemoteRuntimeState() + rcopy := msh.GetRemoteCopy() + displayName := rstate.GetDisplayName(rptr) + rtn := &ResolvedRemote{ + DisplayName: displayName, + RemotePtr: *rptr, + RState: rstate, + MShell: msh, + RemoteCopy: &rcopy, + StatePtr: nil, + FeState: nil, + } + if sessionId != "" && screenId != "" { + ri, err := sstore.GetRemoteInstance(ctx, sessionId, screenId, *rptr) + if err != nil { + log.Printf("ERROR resolving remote state '%s': %v\n", displayName, err) + // continue with state set to nil + } else { + if ri == nil { + rtn.StatePtr = msh.GetDefaultStatePtr() + rtn.FeState = msh.GetDefaultFeState() + } else { + rtn.StatePtr = &sstore.ShellStatePtr{BaseHash: ri.StateBaseHash, DiffHashArr: ri.StateDiffHashArr} + rtn.FeState = ri.FeState + } + } + } + return rtn, nil +} + +// returns (remoteDisplayName, remoteptr, state, rstate, err) +func resolveRemote(ctx context.Context, fullRemoteRef string, sessionId string, screenId string) (string, *sstore.RemotePtrType, *remote.RemoteRuntimeState, error) { + if fullRemoteRef == "" { + return "", nil, nil, nil + } + userRef, remoteRef, remoteName, err := parseFullRemoteRef(fullRemoteRef) + if err != nil { + return "", nil, nil, err + } + if userRef != "" { + return "", nil, nil, fmt.Errorf("invalid remote '%s', cannot resolve remote userid '%s'", fullRemoteRef, userRef) + } + rstate := remote.ResolveRemoteRef(remoteRef) + if rstate == nil { + return "", nil, nil, fmt.Errorf("cannot resolve remote '%s': not found", fullRemoteRef) + } + rptr := sstore.RemotePtrType{RemoteId: rstate.RemoteId, Name: remoteName} + rname := rstate.RemoteCanonicalName + if rstate.RemoteAlias != "" { + rname = rstate.RemoteAlias + } + if rptr.Name != "" { + rname = fmt.Sprintf("%s:%s", rname, rptr.Name) + } + return rname, &rptr, rstate, nil +} diff --git a/wavesrv/pkg/cmdrunner/shparse.go b/wavesrv/pkg/cmdrunner/shparse.go new file mode 100644 index 00000000..6b7e4ed7 --- /dev/null +++ b/wavesrv/pkg/cmdrunner/shparse.go @@ -0,0 +1,333 @@ +package cmdrunner + +import ( + "context" + "fmt" + "regexp" + "strings" + + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/apishell/pkg/simpleexpand" + "github.com/commandlinedev/prompt-server/pkg/scpacket" + "github.com/commandlinedev/prompt-server/pkg/utilfn" + "mvdan.cc/sh/v3/expand" + "mvdan.cc/sh/v3/syntax" +) + +var ValidMetaCmdRe = regexp.MustCompile("^/([a-z_][a-z0-9_-]*)(?::([a-z][a-z0-9_-]*))?$") + +type BareMetaCmdDecl struct { + CmdStr string + MetaCmd string +} + +var BareMetaCmds = []BareMetaCmdDecl{ + BareMetaCmdDecl{"cr", "cr"}, + BareMetaCmdDecl{"connect", "cr"}, + BareMetaCmdDecl{"clear", "clear"}, + BareMetaCmdDecl{"reset", "reset"}, + BareMetaCmdDecl{"codeedit", "codeedit"}, + BareMetaCmdDecl{"codeview", "codeview"}, + BareMetaCmdDecl{"imageview", "imageview"}, + BareMetaCmdDecl{"markdownview", "markdownview"}, + BareMetaCmdDecl{"mdview", "markdownview"}, + BareMetaCmdDecl{"csvview", "csvview"}, +} + +const ( + CmdParseTypePositional = "pos" + CmdParseTypeRaw = "raw" +) + +var CmdParseOverrides map[string]string = map[string]string{ + "setenv": CmdParseTypePositional, + "unset": CmdParseTypePositional, + "set": CmdParseTypePositional, + "run": CmdParseTypeRaw, + "comment": CmdParseTypeRaw, + "chat": CmdParseTypeRaw, +} + +func DumpPacket(pk *scpacket.FeCommandPacketType) { + if pk == nil || pk.MetaCmd == "" { + fmt.Printf("[no metacmd]\n") + return + } + if pk.MetaSubCmd == "" { + fmt.Printf("/%s\n", pk.MetaCmd) + } else { + fmt.Printf("/%s:%s\n", pk.MetaCmd, pk.MetaSubCmd) + } + for _, arg := range pk.Args { + fmt.Printf(" %q\n", arg) + } + for key, val := range pk.Kwargs { + fmt.Printf(" [%s]=%q\n", key, val) + } +} + +func isQuoted(source string, w *syntax.Word) bool { + if w == nil { + return false + } + offset := w.Pos().Offset() + if int(offset) >= len(source) { + return false + } + return source[offset] == '"' || source[offset] == '\'' +} + +func getSourceStr(source string, w *syntax.Word) string { + if w == nil { + return "" + } + offset := w.Pos().Offset() + end := w.End().Offset() + return source[offset:end] +} + +func SubMetaCmd(cmd string) string { + switch cmd { + case "s": + return "screen" + case "r": + return "run" + case "c": + return "comment" + case "e": + return "eval" + case "export": + return "setenv" + case "connection": + return "remote" + default: + return cmd + } +} + +// returns (metaCmd, metaSubCmd, rest) +// if metaCmd is "" then this isn't a valid metacmd string +func parseMetaCmd(origCommandStr string) (string, string, string) { + commandStr := strings.TrimSpace(origCommandStr) + if len(commandStr) < 2 { + return "run", "", origCommandStr + } + fields := strings.SplitN(commandStr, " ", 2) + firstArg := fields[0] + rest := "" + if len(fields) > 1 { + rest = strings.TrimSpace(fields[1]) + } + for _, decl := range BareMetaCmds { + if firstArg == decl.CmdStr { + return decl.MetaCmd, "", rest + } + } + m := ValidMetaCmdRe.FindStringSubmatch(firstArg) + if m == nil { + return "run", "", origCommandStr + } + return SubMetaCmd(m[1]), m[2], rest +} + +func onlyPositionalArgs(metaCmd string, metaSubCmd string) bool { + return (CmdParseOverrides[metaCmd] == CmdParseTypePositional) && metaSubCmd == "" +} + +func onlyRawArgs(metaCmd string, metaSubCmd string) bool { + return CmdParseOverrides[metaCmd] == CmdParseTypeRaw +} + +func setBracketArgs(argMap map[string]string, bracketStr string) error { + bracketStr = strings.TrimSpace(bracketStr) + if bracketStr == "" { + return nil + } + strReader := strings.NewReader(bracketStr) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + var wordErr error + var ectx simpleexpand.SimpleExpandContext // do not set HomeDir (we don't expand ~ in bracket args) + err := parser.Words(strReader, func(w *syntax.Word) bool { + litStr, _ := simpleexpand.SimpleExpandWord(ectx, w, bracketStr) + eqIdx := strings.Index(litStr, "=") + var varName, varVal string + if eqIdx == -1 { + varName = litStr + } else { + varName = litStr[0:eqIdx] + varVal = litStr[eqIdx+1:] + } + if !shexec.IsValidBashIdentifier(varName) { + wordErr = fmt.Errorf("invalid identifier %s in bracket args", utilfn.ShellQuote(varName, true, 20)) + return false + } + if varVal == "" { + varVal = "1" + } + argMap[varName] = varVal + return true + }) + if err != nil { + return err + } + if wordErr != nil { + return wordErr + } + return nil +} + +var literalRtnStateCommands = []string{".", "source", "unset", "cd", "alias", "unalias", "deactivate"} + +func getCallExprLitArg(callExpr *syntax.CallExpr, argNum int) string { + if len(callExpr.Args) <= argNum { + return "" + } + arg := callExpr.Args[argNum] + if len(arg.Parts) == 0 { + return "" + } + lit, ok := arg.Parts[0].(*syntax.Lit) + if !ok { + return "" + } + return lit.Value +} + +// detects: export, declare, ., source, X=1, unset +func IsReturnStateCommand(cmdStr string) bool { + cmdReader := strings.NewReader(cmdStr) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(cmdReader, "cmd") + if err != nil { + return false + } + for _, stmt := range file.Stmts { + if callExpr, ok := stmt.Cmd.(*syntax.CallExpr); ok { + if len(callExpr.Assigns) > 0 && len(callExpr.Args) == 0 { + return true + } + arg0 := getCallExprLitArg(callExpr, 0) + if arg0 != "" && utilfn.ContainsStr(literalRtnStateCommands, arg0) { + return true + } + if arg0 == "git" { + arg1 := getCallExprLitArg(callExpr, 1) + if arg1 == "checkout" || arg1 == "switch" { + return true + } + } + } else if _, ok := stmt.Cmd.(*syntax.DeclClause); ok { + return true + } + } + return false +} + +func EvalBracketArgs(origCmdStr string) (map[string]string, string, error) { + rtn := make(map[string]string) + if strings.HasPrefix(origCmdStr, " ") { + rtn["nohist"] = "1" + } + cmdStr := strings.TrimSpace(origCmdStr) + if !strings.HasPrefix(cmdStr, "[") { + return rtn, origCmdStr, nil + } + rbIdx := strings.Index(cmdStr, "]") + if rbIdx == -1 { + return nil, "", fmt.Errorf("unmatched '[' found in command") + } + bracketStr := cmdStr[1:rbIdx] + restStr := strings.TrimSpace(cmdStr[rbIdx+1:]) + err := setBracketArgs(rtn, bracketStr) + if err != nil { + return nil, "", err + } + return rtn, restStr, nil +} + +func unescapeBackSlashes(s string) string { + if strings.Index(s, "\\") == -1 { + return s + } + var newStr []rune + var lastSlash bool + for _, r := range s { + if lastSlash { + lastSlash = false + newStr = append(newStr, r) + continue + } + if r == '\\' { + lastSlash = true + continue + } + newStr = append(newStr, r) + } + return string(newStr) +} + +func EvalMetaCommand(ctx context.Context, origPk *scpacket.FeCommandPacketType) (*scpacket.FeCommandPacketType, error) { + if len(origPk.Args) == 0 { + return nil, fmt.Errorf("empty command (no fields)") + } + if strings.TrimSpace(origPk.Args[0]) == "" { + return nil, fmt.Errorf("empty command") + } + bracketArgs, cmdStr, err := EvalBracketArgs(origPk.Args[0]) + if err != nil { + return nil, err + } + metaCmd, metaSubCmd, commandArgs := parseMetaCmd(cmdStr) + rtnPk := scpacket.MakeFeCommandPacket() + rtnPk.MetaCmd = metaCmd + rtnPk.MetaSubCmd = metaSubCmd + rtnPk.Kwargs = make(map[string]string) + rtnPk.UIContext = origPk.UIContext + rtnPk.RawStr = origPk.RawStr + for key, val := range origPk.Kwargs { + rtnPk.Kwargs[key] = val + } + for key, val := range bracketArgs { + rtnPk.Kwargs[key] = val + } + if onlyRawArgs(metaCmd, metaSubCmd) { + // don't evaluate arguments for /run or /comment + rtnPk.Args = []string{commandArgs} + return rtnPk, nil + } + commandReader := strings.NewReader(commandArgs) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + var words []*syntax.Word + err = parser.Words(commandReader, func(w *syntax.Word) bool { + words = append(words, w) + return true + }) + if err != nil { + return nil, fmt.Errorf("parsing metacmd, position %v", err) + } + envMap := make(map[string]string) // later we can add vars like session, screen, remote, and user + cfg := shexec.GetParserConfig(envMap) + // process arguments + for idx, w := range words { + literalVal, err := expand.Literal(cfg, w) + if err != nil { + return nil, fmt.Errorf("error evaluating metacmd argument %d [%s]: %v", idx+1, getSourceStr(commandArgs, w), err) + } + if isQuoted(commandArgs, w) || onlyPositionalArgs(metaCmd, metaSubCmd) { + rtnPk.Args = append(rtnPk.Args, literalVal) + continue + } + eqIdx := strings.Index(literalVal, "=") + if eqIdx != -1 && eqIdx != 0 { + varName := literalVal[:eqIdx] + varVal := literalVal[eqIdx+1:] + rtnPk.Kwargs[varName] = varVal + continue + } + rtnPk.Args = append(rtnPk.Args, unescapeBackSlashes(literalVal)) + } + if resolveBool(rtnPk.Kwargs["dump"], false) { + DumpPacket(rtnPk) + } + return rtnPk, nil +} diff --git a/wavesrv/pkg/cmdrunner/shparse_test.go b/wavesrv/pkg/cmdrunner/shparse_test.go new file mode 100644 index 00000000..682d028e --- /dev/null +++ b/wavesrv/pkg/cmdrunner/shparse_test.go @@ -0,0 +1,57 @@ +package cmdrunner + +import ( + "fmt" + "os" + "testing" +) + +func xTestParseAliases(t *testing.T) { + m, err := ParseAliases(` +alias cdg='cd work/gopath/src/github.com/sawka' +alias s='scripthaus' +alias x='ls;ls"' +alias foo="bar \"hello\"" +alias x=y +`) + if err != nil { + fmt.Printf("err: %v\n", err) + return + } + fmt.Printf("m: %#v\n", m) +} + +func xTestParseFuncs(t *testing.T) { + file, err := os.ReadFile("./linux-decls.txt") + if err != nil { + t.Fatalf("error reading linux-decls: %v", err) + } + m, err := ParseFuncs(string(file)) + if err != nil { + t.Fatalf("error parsing funcs: %v", err) + } + for key, val := range m { + fmt.Printf("func: %s %d\n", key, len(val)) + } +} + +func testRSC(t *testing.T, cmd string, expected bool) { + rtn := IsReturnStateCommand(cmd) + if rtn != expected { + t.Errorf("cmd [%s], rtn=%v, expected=%v", cmd, rtn, expected) + } +} + +func TestIsReturnStateCommand(t *testing.T) { + testRSC(t, "FOO=1", true) + testRSC(t, "FOO=1 X=2", true) + testRSC(t, "ls", false) + testRSC(t, "export X", true) + testRSC(t, "export X=1", true) + testRSC(t, "declare -x FOO=1", true) + testRSC(t, "source ./test", true) + testRSC(t, "unset FOO BAR", true) + testRSC(t, "FOO=1; ls", true) + testRSC(t, ". ./test", true) + testRSC(t, "{ FOO=6; }", false) +} diff --git a/wavesrv/pkg/cmdrunner/termopts.go b/wavesrv/pkg/cmdrunner/termopts.go new file mode 100644 index 00000000..563690c3 --- /dev/null +++ b/wavesrv/pkg/cmdrunner/termopts.go @@ -0,0 +1,120 @@ +package cmdrunner + +import ( + "fmt" + "strconv" + "strings" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/sstore" +) + +// PTERM=MxM,Mx25 +// PTERM="Mx25!" +// PTERM=80x25,80x35 + +type PTermOptsType struct { + Rows string + RowsFlex bool + Cols string + ColsFlex bool +} + +const PTermMax = "M" + +func isDigits(s string) bool { + for _, ch := range s { + if ch < '0' || ch > '9' { + return false + } + } + return true +} + +func atoiDefault(s string, def int) int { + ival, err := strconv.Atoi(s) + if err != nil { + return def + } + return ival +} + +func parseTermPart(part string, partType string) (string, bool, error) { + flex := true + if strings.HasSuffix(part, "!") { + part = part[:len(part)-1] + flex = false + } + if part == "" { + return PTermMax, flex, nil + } + if part == PTermMax { + return PTermMax, flex, nil + } + if !isDigits(part) { + return "", false, fmt.Errorf("invalid PTERM %s: must be '%s' or [number]", partType, PTermMax) + } + return part, flex, nil +} + +func parseSingleTermStr(s string) (*PTermOptsType, error) { + s = strings.TrimSpace(s) + xIdx := strings.Index(s, "x") + if xIdx == -1 { + return nil, fmt.Errorf("invalid PTERM, must include 'x' to separate width and height (e.g. WxH)") + } + rowsPart := s[0:xIdx] + colsPart := s[xIdx+1:] + rows, rowsFlex, err := parseTermPart(rowsPart, "rows") + if err != nil { + return nil, err + } + cols, colsFlex, err := parseTermPart(colsPart, "cols") + if err != nil { + return nil, err + } + return &PTermOptsType{Rows: rows, RowsFlex: rowsFlex, Cols: cols, ColsFlex: colsFlex}, nil +} + +func GetUITermOpts(winSize *packet.WinSize, ptermStr string) (*packet.TermOpts, error) { + opts, err := parseSingleTermStr(ptermStr) + if err != nil { + return nil, err + } + termOpts := &packet.TermOpts{Rows: shexec.DefaultTermRows, Cols: shexec.DefaultTermCols, Term: remote.DefaultTerm, MaxPtySize: shexec.DefaultMaxPtySize} + if winSize == nil { + winSize = &packet.WinSize{Rows: shexec.DefaultTermRows, Cols: shexec.DefaultTermCols} + } + if winSize.Rows == 0 { + winSize.Rows = shexec.DefaultTermRows + } + if winSize.Cols == 0 { + winSize.Cols = shexec.DefaultTermCols + } + if opts.Rows == PTermMax { + termOpts.Rows = winSize.Rows + } else { + termOpts.Rows = atoiDefault(opts.Rows, termOpts.Rows) + } + if opts.Cols == PTermMax { + termOpts.Cols = winSize.Cols + } else { + termOpts.Cols = atoiDefault(opts.Cols, termOpts.Cols) + } + termOpts.MaxPtySize = base.BoundInt64(termOpts.MaxPtySize, shexec.MinMaxPtySize, shexec.MaxMaxPtySize) + termOpts.Cols = base.BoundInt(termOpts.Cols, shexec.MinTermCols, shexec.MaxTermCols) + termOpts.Rows = base.BoundInt(termOpts.Rows, shexec.MinTermRows, shexec.MaxTermRows) + return termOpts, nil +} + +func convertTermOpts(pkto *packet.TermOpts) *sstore.TermOpts { + return &sstore.TermOpts{ + Rows: int64(pkto.Rows), + Cols: int64(pkto.Cols), + FlexRows: true, + MaxPtySize: pkto.MaxPtySize, + } +} diff --git a/wavesrv/pkg/comp/comp.go b/wavesrv/pkg/comp/comp.go new file mode 100644 index 00000000..610043ae --- /dev/null +++ b/wavesrv/pkg/comp/comp.go @@ -0,0 +1,640 @@ +// scripthaus completion +package comp + +import ( + "bytes" + "context" + "fmt" + "sort" + "strconv" + "strings" + "unicode" + "unicode/utf8" + + "github.com/commandlinedev/apishell/pkg/simpleexpand" + "github.com/commandlinedev/prompt-server/pkg/shparse" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/commandlinedev/prompt-server/pkg/utilfn" + "mvdan.cc/sh/v3/syntax" +) + +const MaxCompQuoteLen = 5000 + +const ( + // local to simplecomp + CGTypeCommand = "command" + CGTypeFile = "file" + CGTypeDir = "directory" + CGTypeVariable = "variable" + + // implemented in cmdrunner + CGTypeMeta = "metacmd" + CGTypeCommandMeta = "command+meta" + + CGTypeRemote = "remote" + CGTypeRemoteInstance = "remoteinstance" + CGTypeGlobalCmd = "globalcmd" +) + +const ( + QuoteTypeLiteral = "" + QuoteTypeDQ = "\"" + QuoteTypeANSI = "$'" + QuoteTypeSQ = "'" +) + +type CompContext struct { + RemotePtr *sstore.RemotePtrType + Cwd string + ForDisplay bool +} + +type ParsedWord struct { + Offset int + Word *syntax.Word + PartialWord string + Prefix string +} + +type CompPoint struct { + StmtStr string + Words []ParsedWord + CompWord int + CompWordPos int + Prefix string + Suffix string +} + +// directories will have a trailing "/" +type CompEntry struct { + Word string + IsMetaCmd bool +} + +type CompReturn struct { + CompType string + Entries []CompEntry + HasMore bool +} + +var noEscChars []bool +var specialEsc []string + +func init() { + noEscChars = make([]bool, 256) + for ch := 0; ch < 256; ch++ { + if (ch >= '0' && ch <= '9') || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || + ch == '-' || ch == '.' || ch == '/' || ch == ':' || ch == '=' { + noEscChars[byte(ch)] = true + } + } + specialEsc = make([]string, 256) + specialEsc[0x7] = "\\a" + specialEsc[0x8] = "\\b" + specialEsc[0x9] = "\\t" + specialEsc[0xa] = "\\n" + specialEsc[0xb] = "\\v" + specialEsc[0xc] = "\\f" + specialEsc[0xd] = "\\r" + specialEsc[0x1b] = "\\E" +} + +func compQuoteDQString(s string, close bool) string { + var buf bytes.Buffer + buf.WriteByte('"') + for _, ch := range s { + if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { + buf.WriteByte('\\') + buf.WriteRune(ch) + continue + } + buf.WriteRune(ch) + } + if close { + buf.WriteByte('"') + } + return buf.String() +} + +func hasGlob(s string) bool { + var lastExtGlob bool + for _, ch := range s { + if ch == '*' || ch == '?' || ch == '[' || ch == '{' { + return true + } + if ch == '+' || ch == '@' || ch == '!' { + lastExtGlob = true + continue + } + if lastExtGlob && ch == '(' { + return true + } + lastExtGlob = false + } + return false +} + +func writeUtf8Literal(buf *bytes.Buffer, ch rune) { + var runeArr [utf8.UTFMax]byte + buf.WriteString("$'") + barr := runeArr[:] + byteLen := utf8.EncodeRune(barr, ch) + for i := 0; i < byteLen; i++ { + buf.WriteString("\\x") + buf.WriteByte(utilfn.HexDigits[barr[i]/16]) + buf.WriteByte(utilfn.HexDigits[barr[i]%16]) + } + buf.WriteByte('\'') +} + +func compQuoteLiteralString(s string) string { + var buf bytes.Buffer + for idx, ch := range s { + if ch == 0 { + break + } + if idx == 0 && ch == '~' { + buf.WriteRune(ch) + continue + } + if ch > unicode.MaxASCII { + writeUtf8Literal(&buf, ch) + continue + } + var bch = byte(ch) + if noEscChars[bch] { + buf.WriteRune(ch) + continue + } + if specialEsc[bch] != "" { + buf.WriteString(specialEsc[bch]) + continue + } + if !unicode.IsPrint(ch) { + writeUtf8Literal(&buf, ch) + continue + } + buf.WriteByte('\\') + buf.WriteByte(bch) + } + return buf.String() +} + +func compQuoteSQString(s string) string { + var buf bytes.Buffer + for _, ch := range s { + if ch == 0 { + break + } + if ch == '\'' { + buf.WriteString("'\\''") + continue + } + var bch byte + if ch <= unicode.MaxASCII { + bch = byte(ch) + } + if ch > unicode.MaxASCII || !unicode.IsPrint(ch) { + buf.WriteByte('\'') + if bch != 0 && specialEsc[bch] != "" { + buf.WriteString(specialEsc[bch]) + } else { + writeUtf8Literal(&buf, ch) + } + buf.WriteByte('\'') + continue + } + buf.WriteByte(bch) + } + return buf.String() +} + +func compQuoteString(s string, quoteType string, close bool) string { + if quoteType != QuoteTypeANSI && quoteType != QuoteTypeLiteral { + for _, ch := range s { + if ch > unicode.MaxASCII || !unicode.IsPrint(ch) || ch == '!' { + quoteType = QuoteTypeANSI + break + } + if ch == '\'' { + if quoteType == QuoteTypeSQ { + quoteType = QuoteTypeANSI + break + } + } + } + } + if quoteType == QuoteTypeANSI { + rtn := strconv.QuoteToASCII(s) + rtn = "$'" + strings.ReplaceAll(rtn[1:len(rtn)-1], "'", "\\'") + if close { + rtn = rtn + "'" + } + return rtn + } + if quoteType == QuoteTypeLiteral { + return compQuoteLiteralString(s) + } + if quoteType == QuoteTypeSQ { + rtn := utilfn.ShellQuote(s, false, MaxCompQuoteLen) + if len(rtn) > 0 && rtn[0] != '\'' { + rtn = "'" + rtn + "'" + } + if !close { + rtn = rtn[0 : len(rtn)-1] + } + return rtn + } + // QuoteTypeDQ + return compQuoteDQString(s, close) +} + +func (p *CompPoint) wordAsStr(w ParsedWord) string { + if w.Word != nil { + return p.StmtStr[w.Word.Pos().Offset():w.Word.End().Offset()] + } + return w.PartialWord +} + +func (p *CompPoint) simpleExpandWord(w ParsedWord) (string, simpleexpand.SimpleExpandInfo) { + ectx := simpleexpand.SimpleExpandContext{} + if w.Word != nil { + return simpleexpand.SimpleExpandWord(ectx, w.Word, p.StmtStr) + } + return simpleexpand.SimpleExpandPartialWord(ectx, w.PartialWord, false) +} + +func getQuoteTypePref(str string) string { + if strings.HasPrefix(str, QuoteTypeANSI) { + return QuoteTypeANSI + } + if strings.HasPrefix(str, QuoteTypeDQ) { + return QuoteTypeDQ + } + if strings.HasPrefix(str, QuoteTypeSQ) { + return QuoteTypeSQ + } + return QuoteTypeLiteral +} + +func (p *CompPoint) getCompPrefix() (string, simpleexpand.SimpleExpandInfo) { + if p.CompWordPos == 0 { + return "", simpleexpand.SimpleExpandInfo{} + } + pword := p.Words[p.CompWord] + wordStr := p.wordAsStr(pword) + if p.CompWordPos == len(wordStr) { + return p.simpleExpandWord(pword) + } + // TODO we can do better, if p.Word is not nil, we can look for which WordPart + // our pos is in. we can then do a normal word expand on the previous parts + // and a partial on just the current part. this is an uncommon case though + // and has very little upside (even bash does not expand multipart words correctly) + partialWordStr := wordStr[:p.CompWordPos] + return simpleexpand.SimpleExpandPartialWord(simpleexpand.SimpleExpandContext{}, partialWordStr, false) +} + +func (p *CompPoint) extendWord(newWord string, newWordComplete bool) utilfn.StrWithPos { + pword := p.Words[p.CompWord] + wordStr := p.wordAsStr(pword) + quotePref := getQuoteTypePref(wordStr) + needsClose := newWordComplete && (len(wordStr) == p.CompWordPos) + wordSuffix := wordStr[p.CompWordPos:] + newQuotedStr := compQuoteString(newWord, quotePref, needsClose) + if needsClose && wordSuffix == "" && !strings.HasSuffix(newWord, "/") { + newQuotedStr = newQuotedStr + " " + } + newPos := len(newQuotedStr) + return utilfn.StrWithPos{Str: newQuotedStr + wordSuffix, Pos: newPos} +} + +// returns (extension, complete) +func computeCompExtension(compPrefix string, crtn *CompReturn) (string, bool) { + if crtn == nil || crtn.HasMore { + return "", false + } + compStrs := crtn.GetCompStrs() + lcp := utilfn.LongestPrefix(compPrefix, compStrs) + if lcp == compPrefix || len(lcp) < len(compPrefix) || !strings.HasPrefix(lcp, compPrefix) { + return "", false + } + return lcp[len(compPrefix):], (utilfn.ContainsStr(compStrs, lcp) && !utilfn.IsPrefix(compStrs, lcp)) +} + +func (p *CompPoint) FullyExtend(crtn *CompReturn) utilfn.StrWithPos { + if crtn == nil || crtn.HasMore { + return utilfn.StrWithPos{Str: p.getOrigStr(), Pos: p.getOrigPos()} + } + compStrs := crtn.GetCompStrs() + compPrefix, _ := p.getCompPrefix() + lcp := utilfn.LongestPrefix(compPrefix, compStrs) + if lcp == compPrefix || len(lcp) < len(compPrefix) || !strings.HasPrefix(lcp, compPrefix) { + return utilfn.StrWithPos{Str: p.getOrigStr(), Pos: p.getOrigPos()} + } + newStr := p.extendWord(lcp, utilfn.ContainsStr(compStrs, lcp)) + var buf bytes.Buffer + buf.WriteString(p.Prefix) + for idx, w := range p.Words { + if idx == p.CompWord { + buf.WriteString(w.Prefix) + buf.WriteString(newStr.Str) + } else { + buf.WriteString(w.Prefix) + buf.WriteString(p.wordAsStr(w)) + } + } + buf.WriteString(p.Suffix) + compWord := p.Words[p.CompWord] + newPos := len(p.Prefix) + compWord.Offset + len(compWord.Prefix) + newStr.Pos + return utilfn.StrWithPos{Str: buf.String(), Pos: newPos} +} + +func (p *CompPoint) dump() { + if p.Prefix != "" { + fmt.Printf("prefix: %s\n", p.Prefix) + } + fmt.Printf("cpos: %d %d\n", p.CompWord, p.CompWordPos) + for idx, w := range p.Words { + fmt.Printf("w[%d]: ", idx) + if w.Prefix != "" { + fmt.Printf("{%s}", w.Prefix) + } + if idx == p.CompWord { + fmt.Printf("%s\n", utilfn.StrWithPos{Str: p.wordAsStr(w), Pos: p.CompWordPos}) + } else { + fmt.Printf("%s\n", p.wordAsStr(w)) + } + } + if p.Suffix != "" { + fmt.Printf("suffix: %s\n", p.Suffix) + } + fmt.Printf("\n") +} + +var SimpleCompGenFns map[string]SimpleCompGenFnType + +func isWhitespace(str string) bool { + return strings.TrimSpace(str) == "" +} + +func splitInitialWhitespace(str string) (string, string) { + for pos, ch := range str { // rune iteration :/ + if !unicode.IsSpace(ch) { + return str[:pos], str[pos:] + } + } + return str, "" +} + +func ParseCompPoint(cmdStr utilfn.StrWithPos) *CompPoint { + fullCmdStr := cmdStr.Str + pos := cmdStr.Pos + // fmt.Printf("---\n") + // fmt.Printf("cmd: %s\n", strWithCursor(fullCmdStr, pos)) + + // first, find the stmt that the pos appears in + cmdReader := strings.NewReader(fullCmdStr) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + var foundStmt *syntax.Stmt + var lastStmt *syntax.Stmt + var restStartPos int + parser.Stmts(cmdReader, func(stmt *syntax.Stmt) bool { // ignore parse errors (since stmtStr will be the unparsed part) + restStartPos = int(stmt.End().Offset()) + lastStmt = stmt + if uint(pos) >= stmt.Pos().Offset() && uint(pos) < stmt.End().Offset() { + foundStmt = stmt + return false + } + // fmt.Printf("stmt: [[%s]] %d:%d (%d)\n", fullCmdStr[stmt.Pos().Offset():stmt.End().Offset()], stmt.Pos().Offset(), stmt.End().Offset(), stmt.Semicolon.Offset()) + return true + }) + restStr := fullCmdStr[restStartPos:] + if foundStmt == nil && lastStmt != nil && isWhitespace(restStr) && lastStmt.Semicolon.Offset() == 0 { + foundStmt = lastStmt + } + var rtnPoint CompPoint + var stmtStr string + var stmtPos int + if foundStmt != nil { + stmtPos = pos - int(foundStmt.Pos().Offset()) + rtnPoint.Prefix = fullCmdStr[:foundStmt.Pos().Offset()] + if isWhitespace(fullCmdStr[foundStmt.End().Offset():]) { + stmtStr = fullCmdStr[foundStmt.Pos().Offset():] + rtnPoint.Suffix = "" + } else { + stmtStr = fullCmdStr[foundStmt.Pos().Offset():foundStmt.End().Offset()] + rtnPoint.Suffix = fullCmdStr[foundStmt.End().Offset():] + } + } else { + stmtStr = restStr + stmtPos = pos - restStartPos + rtnPoint.Prefix = fullCmdStr[:restStartPos] + rtnPoint.Suffix = fullCmdStr[restStartPos+len(stmtStr):] + } + if stmtPos > len(stmtStr) { + // this should not happen and will cause a jump in completed strings + stmtPos = len(stmtStr) + } + // fmt.Printf("found: ((%s))%s((%s))\n", rtnPoint.Prefix, strWithCursor(stmtStr, stmtPos), rtnPoint.Suffix) + + // now, find the word that the pos appears in within the stmt above + rtnPoint.StmtStr = stmtStr + stmtReader := strings.NewReader(stmtStr) + lastWordPos := 0 + parser.Words(stmtReader, func(w *syntax.Word) bool { + var pword ParsedWord + pword.Offset = lastWordPos + if int(w.Pos().Offset()) > lastWordPos { + pword.Prefix = stmtStr[lastWordPos:w.Pos().Offset()] + } + pword.Word = w + rtnPoint.Words = append(rtnPoint.Words, pword) + lastWordPos = int(w.End().Offset()) + return true + }) + if lastWordPos < len(stmtStr) { + pword := ParsedWord{Offset: lastWordPos} + pword.Prefix, pword.PartialWord = splitInitialWhitespace(stmtStr[lastWordPos:]) + rtnPoint.Words = append(rtnPoint.Words, pword) + } + if len(rtnPoint.Words) == 0 { + rtnPoint.Words = append(rtnPoint.Words, ParsedWord{}) + } + for idx, w := range rtnPoint.Words { + wordLen := len(rtnPoint.wordAsStr(w)) + if stmtPos > w.Offset && stmtPos <= w.Offset+len(w.Prefix)+wordLen { + rtnPoint.CompWord = idx + rtnPoint.CompWordPos = stmtPos - w.Offset - len(w.Prefix) + if rtnPoint.CompWordPos < 0 { + splitCompWord(&rtnPoint) + } + } + } + return &rtnPoint +} + +func splitCompWord(p *CompPoint) { + w := p.Words[p.CompWord] + prefixPos := p.CompWordPos + len(w.Prefix) + + w1 := ParsedWord{Offset: w.Offset, Prefix: w.Prefix[:prefixPos]} + w2 := ParsedWord{Offset: w.Offset + prefixPos, Prefix: w.Prefix[prefixPos:], Word: w.Word, PartialWord: w.PartialWord} + p.CompWord = p.CompWord // the same (w1) + p.CompWordPos = 0 // will be at 0 since w1 has a word length of 0 + var newWords []ParsedWord + if p.CompWord > 0 { + newWords = append(newWords, p.Words[0:p.CompWord]...) + } + newWords = append(newWords, w1, w2) + newWords = append(newWords, p.Words[p.CompWord+1:]...) + p.Words = newWords +} + +func getCompType(compPos shparse.CompletionPos) string { + switch compPos.CompType { + case shparse.CompTypeCommandMeta: + return CGTypeCommandMeta + + case shparse.CompTypeCommand: + return CGTypeCommand + + case shparse.CompTypeVar: + return CGTypeVariable + + case shparse.CompTypeArg, shparse.CompTypeBasic, shparse.CompTypeAssignment: + return CGTypeFile + + default: + return CGTypeFile + } +} + +func fixupVarPrefix(varPrefix string) string { + if strings.HasPrefix(varPrefix, "${") { + varPrefix = varPrefix[2:] + if strings.HasSuffix(varPrefix, "}") { + varPrefix = varPrefix[:len(varPrefix)-1] + } + } else if strings.HasPrefix(varPrefix, "$") { + varPrefix = varPrefix[1:] + } + return varPrefix +} + +func DoCompGen(ctx context.Context, cmdStr utilfn.StrWithPos, compCtx CompContext) (*CompReturn, *utilfn.StrWithPos, error) { + words := shparse.Tokenize(cmdStr.Str) + cmds := shparse.ParseCommands(words) + compPos := shparse.FindCompletionPos(cmds, cmdStr.Pos) + if compPos.CompType == shparse.CompTypeInvalid { + return nil, nil, nil + } + var compPrefix string + if compPos.CompWord != nil { + var info shparse.ExpandInfo + compPrefix, info = shparse.SimpleExpandPrefix(shparse.ExpandContext{}, compPos.CompWord, compPos.CompWordOffset) + if info.HasGlob || info.HasExtGlob || info.HasHistory || info.HasSpecial { + return nil, nil, nil + } + if compPos.CompType != shparse.CompTypeVar && info.HasVar { + return nil, nil, nil + } + if compPos.CompType == shparse.CompTypeVar { + compPrefix = fixupVarPrefix(compPrefix) + } + } + scType := getCompType(compPos) + crtn, err := DoSimpleComp(ctx, scType, compPrefix, compCtx, nil) + if err != nil { + return nil, nil, err + } + if compCtx.ForDisplay { + return crtn, nil, nil + } + extensionStr, extensionComplete := computeCompExtension(compPrefix, crtn) + if extensionStr == "" { + return crtn, nil, nil + } + rtnSP := compPos.Extend(cmdStr, extensionStr, extensionComplete) + return crtn, &rtnSP, nil +} + +func DoCompGenOld(ctx context.Context, sp utilfn.StrWithPos, compCtx CompContext) (*CompReturn, *utilfn.StrWithPos, error) { + compPoint := ParseCompPoint(sp) + compType := CGTypeFile + if compPoint.CompWord == 0 { + compType = CGTypeCommandMeta + } + // TODO lookup special types + compPrefix, info := compPoint.getCompPrefix() + if info.HasVar || info.HasGlob || info.HasExtGlob || info.HasHistory || info.HasSpecial { + return nil, nil, nil + } + crtn, err := DoSimpleComp(ctx, compType, compPrefix, compCtx, nil) + if err != nil { + return nil, nil, err + } + if compCtx.ForDisplay { + return crtn, nil, nil + } + rtnSP := compPoint.FullyExtend(crtn) + return crtn, &rtnSP, nil +} + +func SortCompReturnEntries(c *CompReturn) { + sort.Slice(c.Entries, func(i int, j int) bool { + e1 := c.Entries[i] + e2 := c.Entries[j] + if e1.Word < e2.Word { + return true + } + if e1.Word == e2.Word && e1.IsMetaCmd && !e2.IsMetaCmd { + return true + } + return false + }) +} + +func CombineCompReturn(compType string, c1 *CompReturn, c2 *CompReturn) *CompReturn { + if c1 == nil { + return c2 + } + if c2 == nil { + return c1 + } + var rtn CompReturn + rtn.CompType = compType + rtn.HasMore = c1.HasMore || c2.HasMore + rtn.Entries = append([]CompEntry{}, c1.Entries...) + rtn.Entries = append(rtn.Entries, c2.Entries...) + SortCompReturnEntries(&rtn) + return &rtn +} + +func (c *CompReturn) GetCompStrs() []string { + rtn := make([]string, len(c.Entries)) + for idx, entry := range c.Entries { + rtn[idx] = entry.Word + } + return rtn +} + +func (c *CompReturn) GetCompDisplayStrs() []string { + rtn := make([]string, len(c.Entries)) + for idx, entry := range c.Entries { + if entry.IsMetaCmd { + rtn[idx] = "^" + entry.Word + } else { + rtn[idx] = entry.Word + } + } + return rtn +} + +func (p CompPoint) getOrigPos() int { + pword := p.Words[p.CompWord] + return len(p.Prefix) + pword.Offset + len(pword.Prefix) + p.CompWordPos +} + +func (p CompPoint) getOrigStr() string { + return p.Prefix + p.StmtStr + p.Suffix +} diff --git a/wavesrv/pkg/comp/comp_test.go b/wavesrv/pkg/comp/comp_test.go new file mode 100644 index 00000000..31b4ed75 --- /dev/null +++ b/wavesrv/pkg/comp/comp_test.go @@ -0,0 +1,106 @@ +package comp + +import ( + "fmt" + "strings" + "testing" +) + +func parseToSP(s string) StrWithPos { + idx := strings.Index(s, "[*]") + if idx == -1 { + return StrWithPos{Str: s} + } + return StrWithPos{Str: s[0:idx] + s[idx+3:], Pos: idx} +} + +func testParse(cmdStr string, pos int) { + fmt.Printf("cmd: %s\n", strWithCursor(cmdStr, pos)) + p := ParseCompPoint(StrWithPos{Str: cmdStr, Pos: pos}) + p.dump() +} + +func _Test1(t *testing.T) { + testParse("ls ", 3) + testParse("ls ", 4) + testParse("ls -l foo", 4) + testParse("ls foo; cd h", 12) + testParse("ls foo; cd h;", 13) + testParse("ls & foo; cd h", 12) + testParse("ls \"he", 6) + testParse("ls;", 3) + testParse("ls;", 2) + testParse("ls; cd x; ls", 8) + testParse("cd \"foo ", 8) + testParse("ls; { ls f", 10) + testParse("ls; { ls -l; ls f", 17) + testParse("ls $(ls f", 9) +} + +func testMiniExtend(t *testing.T, p *CompPoint, newWord string, complete bool, expectedStr string) { + newSP := p.extendWord(newWord, complete) + expectedSP := parseToSP(expectedStr) + if newSP != expectedSP { + t.Fatalf("not equal: [%s] != [%s]", newSP, expectedSP) + } else { + fmt.Printf("extend: %s\n", newSP) + } +} + +func Test2(t *testing.T) { + p := ParseCompPoint(parseToSP("ls f[*]")) + testMiniExtend(t, p, "foo", false, "foo[*]") + testMiniExtend(t, p, "foo", true, "foo [*]") + testMiniExtend(t, p, "foo bar", true, "'foo bar' [*]") + testMiniExtend(t, p, "foo'bar", true, `$'foo\'bar' [*]`) + + p = ParseCompPoint(parseToSP("ls f[*]more")) + testMiniExtend(t, p, "foo", false, "foo[*]more") + testMiniExtend(t, p, "foo bar", false, `'foo bar[*]more`) + testMiniExtend(t, p, "foo bar", true, `'foo bar[*]more`) + testMiniExtend(t, p, "foo's", true, `$'foo\'s[*]more`) +} + +func testParseRT(t *testing.T, origSP StrWithPos) { + p := ParseCompPoint(origSP) + newSP := StrWithPos{Str: p.getOrigStr(), Pos: p.getOrigPos()} + if origSP != newSP { + t.Fatalf("not equal: [%s] != [%s]", origSP, newSP) + } +} + +func Test3(t *testing.T) { + testParseRT(t, parseToSP("ls f[*]")) + testParseRT(t, parseToSP("ls f[*]; more $FOO")) + testParseRT(t, parseToSP("hello; ls [*]f")) + testParseRT(t, parseToSP("ls -l; ./foo he[*]ll more; touch foo &")) +} + +func testExtend(t *testing.T, origStr string, compStrs []string, expectedStr string) { + origSP := parseToSP(origStr) + expectedSP := parseToSP(expectedStr) + p := ParseCompPoint(origSP) + crtn := compsToCompReturn(compStrs, false) + newSP := p.FullyExtend(crtn) + if newSP != expectedSP { + t.Fatalf("comp-fail: %s + %v => [%s] expected[%s]", origSP, compStrs, newSP, expectedSP) + } else { + fmt.Printf("comp: %s + %v => [%s]\n", origSP, compStrs, newSP) + } +} + +func Test4(t *testing.T) { + testExtend(t, "ls f[*]", []string{"foo"}, "ls foo [*]") + testExtend(t, "ls f[*]", []string{"foox", "fooy"}, "ls foo[*]") + testExtend(t, "w; ls f[*]; touch x", []string{"foo"}, "w; ls foo [*]; touch x") + testExtend(t, "w; ls f[*] more; touch x", []string{"foo"}, "w; ls foo [*] more; touch x") + testExtend(t, "w; ls f[*]oo; touch x", []string{"foo"}, "w; ls foo[*]oo; touch x") + testExtend(t, `ls "f[*]`, []string{"foo"}, `ls "foo" [*]`) + testExtend(t, `ls 'f[*]`, []string{"foo"}, `ls 'foo' [*]`) + testExtend(t, `ls $'f[*]`, []string{"foo"}, `ls $'foo' [*]`) + testExtend(t, `ls f[*]`, []string{"foo/"}, `ls foo/[*]`) + testExtend(t, `ls f[*]`, []string{"foo bar"}, `ls 'foo bar' [*]`) + testExtend(t, `ls f[*]`, []string{"f\x01\x02"}, `ls $'f\x01\x02' [*]`) + testExtend(t, `ls "foo [*]`, []string{"foo bar"}, `ls "foo bar" [*]`) + testExtend(t, `ls f[*]`, []string{"foo's"}, `ls $'foo\'s' [*]`) +} diff --git a/wavesrv/pkg/comp/simplecomp.go b/wavesrv/pkg/comp/simplecomp.go new file mode 100644 index 00000000..b094eb5f --- /dev/null +++ b/wavesrv/pkg/comp/simplecomp.go @@ -0,0 +1,100 @@ +package comp + +import ( + "context" + "fmt" + "sync" + + "github.com/google/uuid" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/utilfn" +) + +var globalLock = &sync.Mutex{} +var simpleCompMap = map[string]SimpleCompGenFnType{ + CGTypeCommand: simpleCompCommand, + CGTypeFile: simpleCompFile, + CGTypeDir: simpleCompDir, + CGTypeVariable: simpleCompVar, +} + +type SimpleCompGenFnType = func(ctx context.Context, prefix string, compCtx CompContext, args []interface{}) (*CompReturn, error) + +func RegisterSimpleCompFn(compType string, fn SimpleCompGenFnType) { + globalLock.Lock() + defer globalLock.Unlock() + if _, ok := simpleCompMap[compType]; ok { + panic(fmt.Sprintf("simpleCompFn %q already registered", compType)) + } + simpleCompMap[compType] = fn +} + +func getSimpleCompFn(compType string) SimpleCompGenFnType { + globalLock.Lock() + defer globalLock.Unlock() + return simpleCompMap[compType] +} + +func DoSimpleComp(ctx context.Context, compType string, prefix string, compCtx CompContext, args []interface{}) (*CompReturn, error) { + compFn := getSimpleCompFn(compType) + if compFn == nil { + return nil, fmt.Errorf("no simple comp fn for %q", compType) + } + crtn, err := compFn(ctx, prefix, compCtx, args) + if err != nil { + return nil, err + } + crtn.CompType = compType + return crtn, nil +} + +func compsToCompReturn(comps []string, hasMore bool) *CompReturn { + var rtn CompReturn + rtn.HasMore = hasMore + for _, comp := range comps { + rtn.Entries = append(rtn.Entries, CompEntry{Word: comp}) + } + return &rtn +} + +func doCompGen(ctx context.Context, prefix string, compType string, compCtx CompContext) (*CompReturn, error) { + if !packet.IsValidCompGenType(compType) { + return nil, fmt.Errorf("/_compgen invalid type '%s'", compType) + } + msh := remote.GetRemoteById(compCtx.RemotePtr.RemoteId) + if msh == nil { + return nil, fmt.Errorf("invalid remote '%s', not found", compCtx.RemotePtr) + } + cgPacket := packet.MakeCompGenPacket() + cgPacket.ReqId = uuid.New().String() + cgPacket.CompType = compType + cgPacket.Prefix = prefix + cgPacket.Cwd = compCtx.Cwd + resp, err := msh.PacketRpc(ctx, cgPacket) + if err != nil { + return nil, err + } + if err = resp.Err(); err != nil { + return nil, err + } + comps := utilfn.GetStrArr(resp.Data, "comps") + hasMore := utilfn.GetBool(resp.Data, "hasmore") + return compsToCompReturn(comps, hasMore), nil +} + +func simpleCompFile(ctx context.Context, prefix string, compCtx CompContext, args []interface{}) (*CompReturn, error) { + return doCompGen(ctx, prefix, CGTypeFile, compCtx) +} + +func simpleCompDir(ctx context.Context, prefix string, compCtx CompContext, args []interface{}) (*CompReturn, error) { + return doCompGen(ctx, prefix, CGTypeDir, compCtx) +} + +func simpleCompVar(ctx context.Context, prefix string, compCtx CompContext, args []interface{}) (*CompReturn, error) { + return doCompGen(ctx, prefix, CGTypeVariable, compCtx) +} + +func simpleCompCommand(ctx context.Context, prefix string, compCtx CompContext, args []interface{}) (*CompReturn, error) { + return doCompGen(ctx, prefix, CGTypeCommand, compCtx) +} diff --git a/wavesrv/pkg/dbutil/dbutil.go b/wavesrv/pkg/dbutil/dbutil.go new file mode 100644 index 00000000..36387327 --- /dev/null +++ b/wavesrv/pkg/dbutil/dbutil.go @@ -0,0 +1,212 @@ +package dbutil + +import ( + "database/sql/driver" + "encoding/json" + "fmt" + "reflect" + "strconv" +) + +func QuickSetStr(strVal *string, m map[string]interface{}, name string) { + v, ok := m[name] + if !ok { + return + } + ival, ok := v.(int64) + if ok { + *strVal = strconv.FormatInt(ival, 10) + return + } + str, ok := v.(string) + if !ok { + return + } + *strVal = str +} + +func QuickSetInt(ival *int, m map[string]interface{}, name string) { + v, ok := m[name] + if !ok { + return + } + sqlInt, ok := v.(int) + if ok { + *ival = sqlInt + return + } + sqlInt64, ok := v.(int64) + if ok { + *ival = int(sqlInt64) + return + } +} + +func QuickSetInt64(ival *int64, m map[string]interface{}, name string) { + v, ok := m[name] + if !ok { + return + } + sqlInt64, ok := v.(int64) + if ok { + *ival = sqlInt64 + return + } + sqlInt, ok := v.(int) + if ok { + *ival = int64(sqlInt) + return + } +} + +func QuickSetBool(bval *bool, m map[string]interface{}, name string) { + v, ok := m[name] + if !ok { + return + } + sqlInt, ok := v.(int64) + if ok { + if sqlInt > 0 { + *bval = true + } + return + } + sqlBool, ok := v.(bool) + if ok { + *bval = sqlBool + } +} + +func QuickSetBytes(bval *[]byte, m map[string]interface{}, name string) { + v, ok := m[name] + if !ok { + return + } + sqlBytes, ok := v.([]byte) + if ok { + *bval = sqlBytes + } +} + +func getByteArr(m map[string]any, name string, def string) ([]byte, bool) { + v, ok := m[name] + if !ok { + return nil, false + } + barr, ok := v.([]byte) + if !ok { + str, ok := v.(string) + if !ok { + return nil, false + } + barr = []byte(str) + } + if len(barr) == 0 { + barr = []byte(def) + } + return barr, true +} + +func QuickSetJson(ptr interface{}, m map[string]interface{}, name string) { + barr, ok := getByteArr(m, name, "{}") + if !ok { + return + } + json.Unmarshal(barr, ptr) +} + +func QuickSetNullableJson(ptr interface{}, m map[string]interface{}, name string) { + barr, ok := getByteArr(m, name, "null") + if !ok { + return + } + json.Unmarshal(barr, ptr) +} + +func QuickSetJsonArr(ptr interface{}, m map[string]interface{}, name string) { + barr, ok := getByteArr(m, name, "[]") + if !ok { + return + } + json.Unmarshal(barr, ptr) +} + +func CheckNil(v interface{}) bool { + rv := reflect.ValueOf(v) + if !rv.IsValid() { + return true + } + switch rv.Kind() { + case reflect.Chan, reflect.Func, reflect.Interface, reflect.Map, reflect.Pointer, reflect.Slice: + return rv.IsNil() + + default: + return false + } +} + +func QuickNullableJson(v interface{}) string { + if CheckNil(v) { + return "null" + } + barr, _ := json.Marshal(v) + return string(barr) +} + +func QuickJson(v interface{}) string { + if CheckNil(v) { + return "{}" + } + barr, _ := json.Marshal(v) + return string(barr) +} + +func QuickJsonBytes(v interface{}) []byte { + if CheckNil(v) { + return []byte("{}") + } + barr, _ := json.Marshal(v) + return barr +} + +func QuickJsonArr(v interface{}) string { + if CheckNil(v) { + return "[]" + } + barr, _ := json.Marshal(v) + return string(barr) +} + +func QuickJsonArrBytes(v interface{}) []byte { + if CheckNil(v) { + return []byte("[]") + } + barr, _ := json.Marshal(v) + return barr +} + +func QuickScanJson(ptr interface{}, val interface{}) error { + barrVal, ok := val.([]byte) + if !ok { + strVal, ok := val.(string) + if !ok { + return fmt.Errorf("cannot scan '%T' into '%T'", val, ptr) + } + barrVal = []byte(strVal) + } + if len(barrVal) == 0 { + barrVal = []byte("{}") + } + return json.Unmarshal(barrVal, ptr) +} + +func QuickValueJson(v interface{}) (driver.Value, error) { + if CheckNil(v) { + return "{}", nil + } + barr, err := json.Marshal(v) + if err != nil { + return nil, err + } + return string(barr), nil +} diff --git a/wavesrv/pkg/dbutil/map.go b/wavesrv/pkg/dbutil/map.go new file mode 100644 index 00000000..dffe922b --- /dev/null +++ b/wavesrv/pkg/dbutil/map.go @@ -0,0 +1,234 @@ +package dbutil + +import ( + "fmt" + "reflect" + "strings" + + "github.com/sawka/txwrap" +) + +type DBMappable interface { + UseDBMap() +} + +type MapEntry[T any] struct { + Key string + Val T +} + +type MapConverter interface { + ToMap() map[string]interface{} + FromMap(map[string]interface{}) bool +} + +type HasSimpleKey interface { + GetSimpleKey() string +} + +type HasSimpleInt64Key interface { + GetSimpleKey() int64 +} + +type MapConverterPtr[T any] interface { + MapConverter + *T +} + +type DBMappablePtr[T any] interface { + DBMappable + *T +} + +func FromMap[PT MapConverterPtr[T], T any](m map[string]any) PT { + if len(m) == 0 { + return nil + } + rtn := PT(new(T)) + ok := rtn.FromMap(m) + if !ok { + return nil + } + return rtn +} + +func GetMapGen[PT MapConverterPtr[T], T any](tx *txwrap.TxWrap, query string, args ...interface{}) PT { + m := tx.GetMap(query, args...) + return FromMap[PT](m) +} + +func GetMappable[PT DBMappablePtr[T], T any](tx *txwrap.TxWrap, query string, args ...interface{}) PT { + m := tx.GetMap(query, args...) + if len(m) == 0 { + return nil + } + rtn := PT(new(T)) + FromDBMap(rtn, m) + return rtn +} + +func SelectMappable[PT DBMappablePtr[T], T any](tx *txwrap.TxWrap, query string, args ...interface{}) []PT { + var rtn []PT + marr := tx.SelectMaps(query, args...) + for _, m := range marr { + if len(m) == 0 { + continue + } + val := PT(new(T)) + FromDBMap(val, m) + rtn = append(rtn, val) + } + return rtn +} + +func SelectMapsGen[PT MapConverterPtr[T], T any](tx *txwrap.TxWrap, query string, args ...interface{}) []PT { + var rtn []PT + marr := tx.SelectMaps(query, args...) + for _, m := range marr { + val := FromMap[PT](m) + if val != nil { + rtn = append(rtn, val) + } + } + return rtn +} + +func SelectSimpleMap[T any](tx *txwrap.TxWrap, query string, args ...interface{}) map[string]T { + var rtn []MapEntry[T] + tx.Select(&rtn, query, args...) + if len(rtn) == 0 { + return nil + } + rtnMap := make(map[string]T) + for _, entry := range rtn { + rtnMap[entry.Key] = entry.Val + } + return rtnMap +} + +func MakeGenMap[T HasSimpleKey](arr []T) map[string]T { + rtn := make(map[string]T) + for _, val := range arr { + rtn[val.GetSimpleKey()] = val + } + return rtn +} + +func MakeGenMapInt64[T HasSimpleInt64Key](arr []T) map[int64]T { + rtn := make(map[int64]T) + for _, val := range arr { + rtn[val.GetSimpleKey()] = val + } + return rtn +} + +func isStructType(rt reflect.Type) bool { + if rt.Kind() == reflect.Struct { + return true + } + if rt.Kind() == reflect.Pointer && rt.Elem().Kind() == reflect.Struct { + return true + } + return false +} + +func isByteArrayType(t reflect.Type) bool { + return t.Kind() == reflect.Slice && t.Elem().Kind() == reflect.Uint8 +} + +func isStringMapType(t reflect.Type) bool { + return t.Kind() == reflect.Map && t.Key().Kind() == reflect.String +} + +func ToDBMap(v DBMappable, useBytes bool) map[string]interface{} { + if CheckNil(v) { + return nil + } + rv := reflect.ValueOf(v) + if rv.Kind() == reflect.Pointer { + rv = rv.Elem() + } + if rv.Kind() != reflect.Struct { + panic(fmt.Sprintf("invalid type %T (non-struct) passed to StructToDBMap", v)) + } + rt := rv.Type() + m := make(map[string]interface{}) + numFields := rt.NumField() + for i := 0; i < numFields; i++ { + field := rt.Field(i) + fieldVal := rv.FieldByIndex(field.Index) + dbName := field.Tag.Get("dbmap") + if dbName == "" { + dbName = strings.ToLower(field.Name) + } + if dbName == "-" { + continue + } + if isByteArrayType(field.Type) { + m[dbName] = fieldVal.Interface() + } else if field.Type.Kind() == reflect.Slice { + if useBytes { + m[dbName] = QuickJsonArrBytes(fieldVal.Interface()) + } else { + m[dbName] = QuickJsonArr(fieldVal.Interface()) + } + } else if isStructType(field.Type) || isStringMapType(field.Type) { + if useBytes { + m[dbName] = QuickJsonBytes(fieldVal.Interface()) + } else { + m[dbName] = QuickJson(fieldVal.Interface()) + } + } else { + m[dbName] = fieldVal.Interface() + } + } + return m +} + +func FromDBMap(v DBMappable, m map[string]interface{}) { + if CheckNil(v) { + panic("StructFromDBMap, v cannot be nil") + } + rv := reflect.ValueOf(v) + if rv.Kind() == reflect.Pointer { + rv = rv.Elem() + } + if rv.Kind() != reflect.Struct { + panic(fmt.Sprintf("invalid type %T (non-struct) passed to StructFromDBMap", v)) + } + rt := rv.Type() + numFields := rt.NumField() + for i := 0; i < numFields; i++ { + field := rt.Field(i) + fieldVal := rv.FieldByIndex(field.Index) + dbName := field.Tag.Get("dbmap") + if dbName == "" { + dbName = strings.ToLower(field.Name) + } + if dbName == "-" { + continue + } + if isByteArrayType(field.Type) { + barrVal := fieldVal.Addr().Interface() + QuickSetBytes(barrVal.(*[]byte), m, dbName) + } else if field.Type.Kind() == reflect.Slice { + QuickSetJsonArr(fieldVal.Addr().Interface(), m, dbName) + } else if isStructType(field.Type) || isStringMapType(field.Type) { + QuickSetJson(fieldVal.Addr().Interface(), m, dbName) + } else if field.Type.Kind() == reflect.String { + strVal := fieldVal.Addr().Interface() + QuickSetStr(strVal.(*string), m, dbName) + } else if field.Type.Kind() == reflect.Int64 { + intVal := fieldVal.Addr().Interface() + QuickSetInt64(intVal.(*int64), m, dbName) + } else if field.Type.Kind() == reflect.Int { + intVal := fieldVal.Addr().Interface() + QuickSetInt(intVal.(*int), m, dbName) + } else if field.Type.Kind() == reflect.Bool { + boolVal := fieldVal.Addr().Interface() + QuickSetBool(boolVal.(*bool), m, dbName) + } else { + panic(fmt.Sprintf("StructFromDBMap invalid field type %v in %T", fieldVal.Type(), v)) + } + } +} diff --git a/wavesrv/pkg/keygen/keygen.go b/wavesrv/pkg/keygen/keygen.go new file mode 100644 index 00000000..a8ae232e --- /dev/null +++ b/wavesrv/pkg/keygen/keygen.go @@ -0,0 +1,112 @@ +// Utility functions for generating and reading public/private keypairs. +package keygen + +import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "encoding/base64" + "encoding/pem" + "fmt" + "math/big" + "os" + "time" +) + +const p384Params = "BgUrgQQAIg==" + +// Creates a keypair with CN=[id], private key at keyFileName, and +// public key certificate at certFileName. +func CreateKeyPair(keyFileName string, certFileName string, id string) error { + privateKey, err := CreatePrivateKey(keyFileName) + if err != nil { + return err + } + err = CreateCertificate(certFileName, privateKey, id) + if err != nil { + return err + } + return nil +} + +// Creates a private key at keyFileName (ECDSA, secp384r1 (P-384)), PEM format +func CreatePrivateKey() (*ecdsa.PrivateKey, error) { + curve := elliptic.P384() // secp384r1 + privateKey, err := ecdsa.GenerateKey(curve, rand.Reader) + if err != nil { + return nil, fmt.Errorf("Error generating P-384 key err:%w", err) + } + keyFile, err := os.Create(keyFileName) + if err != nil { + return nil, fmt.Errorf("error opening file:%s err:%w", keyFileName, err) + } + defer keyFile.Close() + pkBytes, err := x509.MarshalPKCS8PrivateKey(privateKey) + if err != nil { + return nil, fmt.Errorf("Error MarshalPKCS8PrivateKey err:%w", err) + } + paramsBytes, err := base64.StdEncoding.DecodeString(p384Params) + if err != nil { + return nil, fmt.Errorf("Error decoding bytes for P-384 EC PARAMETERS err:%w", err) + } + var pemParamsBlock = &pem.Block{ + Type: "EC PARAMETERS", + Bytes: paramsBytes, + } + err = pem.Encode(keyFile, pemParamsBlock) + if err != nil { + return nil, fmt.Errorf("Error writing EC PARAMETERS pem block err:%w", err) + } + var pemPrivateBlock = &pem.Block{ + Type: "EC PRIVATE KEY", + Bytes: pkBytes, + } + err = pem.Encode(keyFile, pemPrivateBlock) + if err != nil { + return nil, fmt.Errorf("Error writing EC PRIVATE KEY pem block err:%w", err) + } + return privateKey, nil +} + +// Creates a public key certificate at certFileName using privateKey with CN=[id]. +func CreateCertificate(certFileName string, privateKey *ecdsa.PrivateKey, id string) error { + serialNumber, err := rand.Int(rand.Reader, big.NewInt(1000000000000)) + if err != nil { + return fmt.Errorf("Cannot generate serial number err:%w", err) + } + notBefore, err := time.Parse("Jan 2 15:04:05 2006", "Jan 1 00:00:00 2020") + if err != nil { + return fmt.Errorf("Cannot Parse Date err:%w", err) + } + notAfter, err := time.Parse("Jan 2 15:04:05 2006", "Jan 1 00:00:00 2030") + if err != nil { + return fmt.Errorf("Cannot Parse Date err:%w", err) + } + template := x509.Certificate{ + SerialNumber: serialNumber, + Subject: pkix.Name{ + CommonName: id, + }, + NotBefore: notBefore, + NotAfter: notAfter, + KeyUsage: x509.KeyUsageDigitalSignature, + ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageClientAuth}, + BasicConstraintsValid: true, + } + certBytes, err := x509.CreateCertificate(rand.Reader, &template, &template, &privateKey.PublicKey, privateKey) + if err != nil { + return fmt.Errorf("Error running x509.CreateCertificate err:%v\n", err) + } + certFile, err := os.Create(certFileName) + if err != nil { + return fmt.Errorf("Error opening file:%s err:%w", certFileName, err) + } + defer certFile.Close() + err = pem.Encode(certFile, &pem.Block{Type: "CERTIFICATE", Bytes: certBytes}) + if err != nil { + return fmt.Errorf("Error writing CERTIFICATE pem block err:%w", err) + } + return nil +} diff --git a/wavesrv/pkg/mapqueue/mapqueue.go b/wavesrv/pkg/mapqueue/mapqueue.go new file mode 100644 index 00000000..aa5d4d9a --- /dev/null +++ b/wavesrv/pkg/mapqueue/mapqueue.go @@ -0,0 +1,99 @@ +package mapqueue + +import ( + "fmt" + "log" + "runtime/debug" + "sync" +) + +type MQEntry struct { + Lock *sync.Mutex + Running bool + Queue chan func() +} + +type MapQueue struct { + Lock *sync.Mutex + M map[string]*MQEntry + QueueSize int +} + +func MakeMapQueue(queueSize int) *MapQueue { + rtn := &MapQueue{ + Lock: &sync.Mutex{}, + M: make(map[string]*MQEntry), + QueueSize: queueSize, + } + return rtn +} + +func (mq *MapQueue) getEntry(id string) *MQEntry { + mq.Lock.Lock() + defer mq.Lock.Unlock() + entry := mq.M[id] + if entry == nil { + entry = &MQEntry{ + Lock: &sync.Mutex{}, + Running: false, + Queue: make(chan func(), mq.QueueSize), + } + mq.M[id] = entry + } + return entry +} + +func (entry *MQEntry) add(fn func()) error { + select { + case entry.Queue <- fn: + break + default: + return fmt.Errorf("input queue full") + } + entry.tryRun() + return nil +} + +func runFn(fn func()) { + defer func() { + r := recover() + if r == nil { + return + } + log.Printf("[error] panic in MQEntry runFn: %v\n", r) + debug.PrintStack() + return + }() + fn() +} + +func (entry *MQEntry) tryRun() { + entry.Lock.Lock() + defer entry.Lock.Unlock() + if entry.Running { + return + } + if len(entry.Queue) > 0 { + entry.Running = true + go entry.run() + } +} + +func (entry *MQEntry) run() { + for fn := range entry.Queue { + runFn(fn) + } + entry.Lock.Lock() + entry.Running = false + entry.Lock.Unlock() + entry.tryRun() +} + +func (mq *MapQueue) Enqueue(id string, fn func()) error { + entry := mq.getEntry(id) + err := entry.add(fn) + if err != nil { + return fmt.Errorf("cannot enqueue: %v", err) + } + return nil +} diff --git a/wavesrv/pkg/pcloud/pcloud.go b/wavesrv/pkg/pcloud/pcloud.go new file mode 100644 index 00000000..89b69802 --- /dev/null +++ b/wavesrv/pkg/pcloud/pcloud.go @@ -0,0 +1,628 @@ +package pcloud + +import ( + "bytes" + "context" + "encoding/json" + "errors" + "fmt" + "io" + "log" + "net/http" + "os" + "strconv" + "strings" + "sync" + "time" + + "github.com/commandlinedev/prompt-server/pkg/dbutil" + "github.com/commandlinedev/prompt-server/pkg/rtnstate" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/commandlinedev/prompt-server/pkg/sstore" +) + +const PCloudEndpoint = "https://api.getprompt.dev/central" +const PCloudEndpointVarName = "PCLOUD_ENDPOINT" +const APIVersion = 1 +const MaxPtyUpdateSize = (128 * 1024) +const MaxUpdatesPerReq = 10 +const MaxUpdatesToDeDup = 1000 +const MaxUpdateWriterErrors = 3 +const PCloudDefaultTimeout = 5 * time.Second +const PCloudWebShareUpdateTimeout = 15 * time.Second + +// setting to 1M to be safe (max is 6M for API-GW + Lambda, but there is base64 encoding and upload time) +// we allow one extra update past this estimated size +const MaxUpdatePayloadSize = 1 * (1024 * 1024) + +const TelemetryUrl = "/telemetry" +const NoTelemetryUrl = "/no-telemetry" +const WebShareUpdateUrl = "/auth/web-share-update" + +var updateWriterLock = &sync.Mutex{} +var updateWriterRunning = false +var updateWriterNumFailures = 0 + +type AuthInfo struct { + UserId string `json:"userid"` + ClientId string `json:"clientid"` + AuthKey string `json:"authkey"` +} + +func GetEndpoint() string { + if !scbase.IsDevMode() { + return PCloudEndpoint + } + endpoint := os.Getenv(PCloudEndpointVarName) + if endpoint == "" || !strings.HasPrefix(endpoint, "https://") { + panic("Invalid PCloud dev endpoint, PCLOUD_ENDPOINT not set or invalid") + } + return endpoint +} + +func makeAuthPostReq(ctx context.Context, apiUrl string, authInfo AuthInfo, data interface{}) (*http.Request, error) { + var dataReader io.Reader + if data != nil { + byteArr, err := json.Marshal(data) + if err != nil { + return nil, fmt.Errorf("error marshaling json for %s request: %v", apiUrl, err) + } + dataReader = bytes.NewReader(byteArr) + } + fullUrl := GetEndpoint() + apiUrl + req, err := http.NewRequestWithContext(ctx, "POST", fullUrl, dataReader) + if err != nil { + return nil, fmt.Errorf("error creating %s request: %v", apiUrl, err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-PromptAPIVersion", strconv.Itoa(APIVersion)) + req.Header.Set("X-PromptAPIUrl", apiUrl) + req.Header.Set("X-PromptUserId", authInfo.UserId) + req.Header.Set("X-PromptClientId", authInfo.ClientId) + req.Header.Set("X-PromptAuthKey", authInfo.AuthKey) + req.Close = true + return req, nil +} + +func makeAnonPostReq(ctx context.Context, apiUrl string, data interface{}) (*http.Request, error) { + var dataReader io.Reader + if data != nil { + byteArr, err := json.Marshal(data) + if err != nil { + return nil, fmt.Errorf("error marshaling json for %s request: %v", apiUrl, err) + } + dataReader = bytes.NewReader(byteArr) + } + fullUrl := GetEndpoint() + apiUrl + req, err := http.NewRequestWithContext(ctx, "POST", fullUrl, dataReader) + if err != nil { + return nil, fmt.Errorf("error creating %s request: %v", apiUrl, err) + } + req.Header.Set("Content-Type", "application/json") + req.Header.Set("X-PromptAPIVersion", strconv.Itoa(APIVersion)) + req.Header.Set("X-PromptAPIUrl", apiUrl) + req.Close = true + return req, nil +} + +func doRequest(req *http.Request, outputObj interface{}) (*http.Response, error) { + apiUrl := req.Header.Get("X-PromptAPIUrl") + log.Printf("[pcloud] sending request %s %v\n", req.Method, req.URL) + resp, err := http.DefaultClient.Do(req) + if err != nil { + return nil, fmt.Errorf("error contacting pcloud %q service: %v", apiUrl, err) + } + defer resp.Body.Close() + bodyBytes, err := io.ReadAll(resp.Body) + if err != nil { + return resp, fmt.Errorf("error reading %q response body: %v", apiUrl, err) + } + if resp.StatusCode != http.StatusOK { + return resp, fmt.Errorf("error contacting pcloud %q service: %s", apiUrl, resp.Status) + } + if outputObj != nil && resp.Header.Get("Content-Type") == "application/json" { + err = json.Unmarshal(bodyBytes, outputObj) + if err != nil { + return resp, fmt.Errorf("error decoding json: %v", err) + } + } + return resp, nil +} + +func SendTelemetry(ctx context.Context, force bool) error { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return fmt.Errorf("cannot retrieve client data: %v", err) + } + if !force && clientData.ClientOpts.NoTelemetry { + return nil + } + activity, err := sstore.GetNonUploadedActivity(ctx) + if err != nil { + return fmt.Errorf("cannot get activity: %v", err) + } + if len(activity) == 0 { + return nil + } + log.Printf("[pcloud] sending telemetry data\n") + dayStr := sstore.GetCurDayStr() + input := TelemetryInputType{UserId: clientData.UserId, ClientId: clientData.ClientId, CurDay: dayStr, Activity: activity} + req, err := makeAnonPostReq(ctx, TelemetryUrl, input) + if err != nil { + return err + } + _, err = doRequest(req, nil) + if err != nil { + return err + } + err = sstore.MarkActivityAsUploaded(ctx, activity) + if err != nil { + return fmt.Errorf("error marking activity as uploaded: %v", err) + } + return nil +} + +func SendNoTelemetryUpdate(ctx context.Context, noTelemetryVal bool) error { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return fmt.Errorf("cannot retrieve client data: %v", err) + } + req, err := makeAnonPostReq(ctx, NoTelemetryUrl, NoTelemetryInputType{ClientId: clientData.ClientId, Value: noTelemetryVal}) + if err != nil { + return err + } + _, err = doRequest(req, nil) + if err != nil { + return err + } + return nil +} + +func getAuthInfo(ctx context.Context) (AuthInfo, error) { + clientData, err := sstore.EnsureClientData(ctx) + if err != nil { + return AuthInfo{}, fmt.Errorf("cannot retrieve client data: %v", err) + } + return AuthInfo{UserId: clientData.UserId, ClientId: clientData.ClientId}, nil +} + +func defaultError(err error, estr string) error { + if err != nil { + return err + } + return errors.New(estr) +} + +func MakeScreenNewUpdate(screen *sstore.ScreenType, webShareOpts sstore.ScreenWebShareOpts) *WebShareUpdateType { + rtn := &WebShareUpdateType{ + ScreenId: screen.ScreenId, + UpdateId: -1, + UpdateType: sstore.UpdateType_ScreenNew, + UpdateTs: time.Now().UnixMilli(), + } + rtn.Screen = &WebShareScreenType{ + ScreenId: screen.ScreenId, + SelectedLine: int(screen.SelectedLine), + ShareName: webShareOpts.ShareName, + ViewKey: webShareOpts.ViewKey, + } + return rtn +} + +func MakeScreenDelUpdate(screen *sstore.ScreenType, screenId string) *WebShareUpdateType { + rtn := &WebShareUpdateType{ + ScreenId: screenId, + UpdateId: -1, + UpdateType: sstore.UpdateType_ScreenDel, + UpdateTs: time.Now().UnixMilli(), + } + return rtn +} + +func makeWebShareUpdate(ctx context.Context, update *sstore.ScreenUpdateType) (*WebShareUpdateType, error) { + rtn := &WebShareUpdateType{ + ScreenId: update.ScreenId, + LineId: update.LineId, + UpdateId: update.UpdateId, + UpdateType: update.UpdateType, + UpdateTs: update.UpdateTs, + } + switch update.UpdateType { + case sstore.UpdateType_ScreenNew: + screen, err := sstore.GetScreenById(ctx, update.ScreenId) + if err != nil || screen == nil { + return nil, fmt.Errorf("error getting screen: %v", defaultError(err, "not found")) + } + rtn.Screen, err = webScreenFromScreen(screen) + if err != nil { + return nil, fmt.Errorf("error converting screen to web-screen: %v", err) + } + + case sstore.UpdateType_ScreenDel: + break + + case sstore.UpdateType_ScreenName, sstore.UpdateType_ScreenSelectedLine: + screen, err := sstore.GetScreenById(ctx, update.ScreenId) + if err != nil { + return nil, fmt.Errorf("error getting screen: %v", err) + } + if screen == nil || screen.WebShareOpts == nil { + return nil, fmt.Errorf("invalid screen, not webshared (makeWebScreenUpdate)") + } + if update.UpdateType == sstore.UpdateType_ScreenName { + rtn.SVal = screen.WebShareOpts.ShareName + } else if update.UpdateType == sstore.UpdateType_ScreenSelectedLine { + rtn.IVal = int64(screen.SelectedLine) + } + + case sstore.UpdateType_LineNew: + line, cmd, err := sstore.GetLineCmdByLineId(ctx, update.ScreenId, update.LineId) + if err != nil || line == nil { + return nil, fmt.Errorf("error getting line/cmd: %v", defaultError(err, "not found")) + } + rtn.Line, err = webLineFromLine(line) + if err != nil { + return nil, fmt.Errorf("error converting line to web-line: %v", err) + } + if cmd != nil { + rtn.Cmd, err = webCmdFromCmd(update.LineId, cmd) + if err != nil { + return nil, fmt.Errorf("error converting cmd to web-cmd: %v", err) + } + } + + case sstore.UpdateType_LineDel: + break + + case sstore.UpdateType_LineRenderer, sstore.UpdateType_LineContentHeight: + line, err := sstore.GetLineById(ctx, update.ScreenId, update.LineId) + if err != nil || line == nil { + return nil, fmt.Errorf("error getting line: %v", defaultError(err, "not found")) + } + if update.UpdateType == sstore.UpdateType_LineRenderer { + rtn.SVal = line.Renderer + } else if update.UpdateType == sstore.UpdateType_LineContentHeight { + rtn.IVal = line.ContentHeight + } + + case sstore.UpdateType_CmdStatus: + _, cmd, err := sstore.GetLineCmdByLineId(ctx, update.ScreenId, update.LineId) + if err != nil || cmd == nil { + return nil, fmt.Errorf("error getting cmd: %v", defaultError(err, "not found")) + } + rtn.SVal = cmd.Status + + case sstore.UpdateType_CmdTermOpts: + _, cmd, err := sstore.GetLineCmdByLineId(ctx, update.ScreenId, update.LineId) + if err != nil || cmd == nil { + return nil, fmt.Errorf("error getting cmd: %v", defaultError(err, "not found")) + } + rtn.TermOpts = &cmd.TermOpts + + case sstore.UpdateType_CmdExitCode, sstore.UpdateType_CmdDurationMs: + _, cmd, err := sstore.GetLineCmdByLineId(ctx, update.ScreenId, update.LineId) + if err != nil || cmd == nil { + return nil, fmt.Errorf("error getting cmd: %v", defaultError(err, "not found")) + } + if update.UpdateType == sstore.UpdateType_CmdExitCode { + rtn.IVal = int64(cmd.ExitCode) + } else if update.UpdateType == sstore.UpdateType_CmdDurationMs { + rtn.IVal = int64(cmd.DurationMs) + } + + case sstore.UpdateType_CmdRtnState: + _, cmd, err := sstore.GetLineCmdByLineId(ctx, update.ScreenId, update.LineId) + if err != nil || cmd == nil { + return nil, fmt.Errorf("error getting cmd: %v", defaultError(err, "not found")) + } + data, err := rtnstate.GetRtnStateDiff(ctx, update.ScreenId, cmd.LineId) + if err != nil { + return nil, fmt.Errorf("cannot compute rtnstate: %v", err) + } + rtn.SVal = string(data) + + case sstore.UpdateType_PtyPos: + ptyPos, err := sstore.GetWebPtyPos(ctx, update.ScreenId, update.LineId) + if err != nil { + return nil, fmt.Errorf("error getting ptypos: %v", err) + } + realOffset, data, err := sstore.ReadPtyOutFile(ctx, update.ScreenId, update.LineId, ptyPos, MaxPtyUpdateSize+1) + if err != nil { + return nil, fmt.Errorf("error getting ptydata: %v", err) + } + if len(data) == 0 { + return nil, nil + } + if len(data) > MaxPtyUpdateSize { + rtn.PtyData = &WebSharePtyData{PtyPos: realOffset, Data: data[0:MaxPtyUpdateSize], Eof: false} + } else { + rtn.PtyData = &WebSharePtyData{PtyPos: realOffset, Data: data, Eof: true} + } + + case sstore.UpdateType_LineState: + // TODO implement! + + default: + return nil, fmt.Errorf("unsupported update type (pcloud/makeWebScreenUpdate): %s\n", update.UpdateType) + } + return rtn, nil +} + +func finalizeWebScreenUpdate(ctx context.Context, webUpdate *WebShareUpdateType) error { + switch webUpdate.UpdateType { + case sstore.UpdateType_PtyPos: + newPos := webUpdate.PtyData.PtyPos + int64(len(webUpdate.PtyData.Data)) + err := sstore.SetWebPtyPos(ctx, webUpdate.ScreenId, webUpdate.LineId, newPos) + if err != nil { + return err + } + + case sstore.UpdateType_LineDel: + err := sstore.DeleteWebPtyPos(ctx, webUpdate.ScreenId, webUpdate.LineId) + if err != nil { + return err + } + } + err := sstore.RemoveScreenUpdate(ctx, webUpdate.UpdateId) + if err != nil { + // this is not great, this *should* never fail and is not easy to recover from + return err + } + return nil +} + +type webShareResponseType struct { + Success bool `json:"success"` + Data []*WebShareUpdateResponseType `json:"data"` +} + +func convertUpdate(update *sstore.ScreenUpdateType) *WebShareUpdateType { + webUpdate, err := makeWebShareUpdate(context.Background(), update) + if err != nil || webUpdate == nil { + if err != nil { + log.Printf("[pcloud] error create web-share update updateid:%d: %v", update.UpdateId, err) + } + // if err, or no web update created, remove the screenupdate + removeErr := sstore.RemoveScreenUpdate(context.Background(), update.UpdateId) + if removeErr != nil { + // ignore this error too (although this is really problematic, there is nothing to do) + log.Printf("[pcloud] error removing screen update updateid:%d: %v", update.UpdateId, removeErr) + } + } + return webUpdate +} + +func DoSyncWebUpdate(webUpdate *WebShareUpdateType) error { + authInfo, err := getAuthInfo(context.Background()) + if err != nil { + return fmt.Errorf("could not get authinfo for request: %v", err) + } + ctx, cancelFn := context.WithTimeout(context.Background(), PCloudDefaultTimeout) + defer cancelFn() + req, err := makeAuthPostReq(ctx, WebShareUpdateUrl, authInfo, []*WebShareUpdateType{webUpdate}) + if err != nil { + return fmt.Errorf("cannot create auth-post-req for %s: %v", WebShareUpdateUrl, err) + } + var resp webShareResponseType + _, err = doRequest(req, &resp) + if err != nil { + return err + } + if len(resp.Data) == 0 { + return fmt.Errorf("invalid response received from server") + } + urt := resp.Data[0] + if urt.Error != "" { + return errors.New(urt.Error) + } + return nil +} + +func DoWebUpdates(webUpdates []*WebShareUpdateType) error { + if len(webUpdates) == 0 { + return nil + } + authInfo, err := getAuthInfo(context.Background()) + if err != nil { + return fmt.Errorf("could not get authinfo for request: %v", err) + } + ctx, cancelFn := context.WithTimeout(context.Background(), PCloudWebShareUpdateTimeout) + defer cancelFn() + req, err := makeAuthPostReq(ctx, WebShareUpdateUrl, authInfo, webUpdates) + if err != nil { + return fmt.Errorf("cannot create auth-post-req for %s: %v", WebShareUpdateUrl, err) + } + var resp webShareResponseType + _, err = doRequest(req, &resp) + if err != nil { + return err + } + respMap := dbutil.MakeGenMapInt64(resp.Data) + for _, update := range webUpdates { + err = finalizeWebScreenUpdate(context.Background(), update) + if err != nil { + // ignore this error (nothing to do) + log.Printf("[pcloud] error finalizing web-update: %v\n", err) + } + resp := respMap[update.UpdateId] + if resp == nil { + resp = &WebShareUpdateResponseType{Success: false, Error: "resp not found"} + } + if resp.Error != "" { + log.Printf("[pcloud] error updateid:%d, type:%s %s/%s err:%v\n", update.UpdateId, update.UpdateType, update.ScreenId, update.LineId, resp.Error) + } + } + return nil +} + +func setUpdateWriterRunning(running bool) { + updateWriterLock.Lock() + defer updateWriterLock.Unlock() + updateWriterRunning = running +} + +func GetUpdateWriterRunning() bool { + updateWriterLock.Lock() + defer updateWriterLock.Unlock() + return updateWriterRunning +} + +func StartUpdateWriter() { + updateWriterLock.Lock() + defer updateWriterLock.Unlock() + if updateWriterRunning { + return + } + updateWriterRunning = true + go runWebShareUpdateWriter() +} + +func computeUpdateWriterBackoff() time.Duration { + updateWriterLock.Lock() + numFailures := updateWriterNumFailures + updateWriterLock.Unlock() + switch numFailures { + case 0: + return 0 + case 1: + return 1 * time.Second + case 2: + return 2 * time.Second + case 3: + return 5 * time.Second + case 4: + return time.Minute + case 5: + return 5 * time.Minute + case 6: + return time.Hour + default: + return time.Hour + } +} + +func incrementUpdateWriterNumFailures() { + updateWriterLock.Lock() + defer updateWriterLock.Unlock() + updateWriterNumFailures++ +} + +func ResetUpdateWriterNumFailures() { + updateWriterLock.Lock() + defer updateWriterLock.Unlock() + updateWriterNumFailures = 0 +} + +func GetUpdateWriterNumFailures() int { + updateWriterLock.Lock() + defer updateWriterLock.Unlock() + return updateWriterNumFailures +} + +type updateKey struct { + ScreenId string + LineId string + UpdateType string +} + +func DeDupUpdates(ctx context.Context, updateArr []*sstore.ScreenUpdateType) ([]*sstore.ScreenUpdateType, error) { + var rtn []*sstore.ScreenUpdateType + var idsToDelete []int64 + umap := make(map[updateKey]bool) + for _, update := range updateArr { + key := updateKey{ScreenId: update.ScreenId, LineId: update.LineId, UpdateType: update.UpdateType} + if umap[key] { + idsToDelete = append(idsToDelete, update.UpdateId) + continue + } + umap[key] = true + rtn = append(rtn, update) + } + if len(idsToDelete) > 0 { + err := sstore.RemoveScreenUpdates(ctx, idsToDelete) + if err != nil { + return nil, fmt.Errorf("error trying to delete screenupdates: %v\n", err) + } + } + return rtn, nil +} + +func runWebShareUpdateWriter() { + defer func() { + setUpdateWriterRunning(false) + }() + log.Printf("[pcloud] starting update writer\n") + numErrors := 0 + for { + if numErrors > MaxUpdateWriterErrors { + log.Printf("[pcloud] update-writer, too many errors, exiting\n") + break + } + time.Sleep(100 * time.Millisecond) + fullUpdateArr, err := sstore.GetScreenUpdates(context.Background(), MaxUpdatesToDeDup) + if err != nil { + log.Printf("[pcloud] error retrieving updates: %v", err) + time.Sleep(1 * time.Second) + numErrors++ + continue + } + updateArr, err := DeDupUpdates(context.Background(), fullUpdateArr) + if err != nil { + log.Printf("[pcloud] error deduping screenupdates: %v", err) + time.Sleep(1 * time.Second) + numErrors++ + continue + } + numErrors = 0 + + var webUpdateArr []*WebShareUpdateType + totalSize := 0 + for _, update := range updateArr { + webUpdate := convertUpdate(update) + if webUpdate == nil { + continue + } + webUpdateArr = append(webUpdateArr, webUpdate) + totalSize += webUpdate.GetEstimatedSize() + if totalSize > MaxUpdatePayloadSize { + break + } + } + if len(webUpdateArr) == 0 { + sstore.UpdateWriterCheckMoreData() + continue + } + err = DoWebUpdates(webUpdateArr) + if err != nil { + incrementUpdateWriterNumFailures() + backoffTime := computeUpdateWriterBackoff() + log.Printf("[pcloud] error processing %d web-updates (backoff=%v): %v\n", len(webUpdateArr), backoffTime, err) + updateBackoffSleep(backoffTime) + continue + } + log.Printf("[pcloud] sent %d web-updates\n", len(webUpdateArr)) + var debugStrs []string + for _, webUpdate := range webUpdateArr { + debugStrs = append(debugStrs, webUpdate.String()) + } + log.Printf("[pcloud] updates: %s\n", strings.Join(debugStrs, " ")) + ResetUpdateWriterNumFailures() + } +} + +// todo fix this, set deadline, check with condition variable, backoff then just needs to notify +func updateBackoffSleep(backoffTime time.Duration) { + var totalSleep time.Duration + for { + sleepTime := time.Second + totalSleep += sleepTime + time.Sleep(sleepTime) + if totalSleep >= backoffTime { + break + } + numFailures := GetUpdateWriterNumFailures() + if numFailures == 0 { + break + } + } +} diff --git a/wavesrv/pkg/pcloud/pclouddata.go b/wavesrv/pkg/pcloud/pclouddata.go new file mode 100644 index 00000000..87e43b4b --- /dev/null +++ b/wavesrv/pkg/pcloud/pclouddata.go @@ -0,0 +1,196 @@ +package pcloud + +import ( + "context" + "encoding/json" + "fmt" + + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/rtnstate" + "github.com/commandlinedev/prompt-server/pkg/sstore" +) + +type NoTelemetryInputType struct { + ClientId string `json:"clientid"` + Value bool `json:"value"` +} + +type TelemetryInputType struct { + UserId string `json:"userid"` + ClientId string `json:"clientid"` + CurDay string `json:"curday"` + Activity []*sstore.ActivityType `json:"activity"` +} + +type WebShareUpdateType struct { + ScreenId string `json:"screenid"` + LineId string `json:"lineid"` + UpdateId int64 `json:"updateid"` + UpdateType string `json:"updatetype"` + UpdateTs int64 `json:"updatets"` + + Screen *WebShareScreenType `json:"screen,omitempty"` + Line *WebShareLineType `json:"line,omitempty"` + Cmd *WebShareCmdType `json:"cmd,omitempty"` + PtyData *WebSharePtyData `json:"ptydata,omitempty"` + SVal string `json:"sval,omitempty"` + IVal int64 `json:"ival,omitempty"` + BVal bool `json:"bval,omitempty"` + TermOpts *sstore.TermOpts `json:"termopts,omitempty"` +} + +const EstimatedSizePadding = 100 + +func (update *WebShareUpdateType) GetEstimatedSize() int { + barr, _ := json.Marshal(update) + return len(barr) + 100 +} + +func (update *WebShareUpdateType) String() string { + var idStr string + if update.LineId != "" && update.ScreenId != "" { + idStr = fmt.Sprintf("%s:%s", update.ScreenId[0:8], update.LineId[0:8]) + } else if update.ScreenId != "" { + idStr = update.ScreenId[0:8] + } + if update.UpdateType == sstore.UpdateType_PtyPos && update.PtyData != nil { + return fmt.Sprintf("ptydata[%s][%d:%d]", idStr, update.PtyData.PtyPos, len(update.PtyData.Data)) + } + return fmt.Sprintf("%s[%s]", update.UpdateType, idStr) +} + +type WebShareUpdateResponseType struct { + UpdateId int64 `json:"updateid"` + Success bool `json:"success"` + Error string `json:"error,omitempty"` +} + +func (ur *WebShareUpdateResponseType) GetSimpleKey() int64 { + return ur.UpdateId +} + +type WebShareRemote struct { + RemoteId string `json:"remoteid"` + Alias string `json:"alias,omitempty"` + CanonicalName string `json:"canonicalname"` + Name string `json:"name,omitempty"` + HomeDir string `json:"homedir,omitempty"` + IsRoot bool `json:"isroot,omitempty"` +} + +type WebShareScreenType struct { + ScreenId string `json:"screenid"` + ShareName string `json:"sharename"` + ViewKey string `json:"viewkey"` + SelectedLine int `json:"selectedline"` +} + +func webRemoteFromRemote(rptr sstore.RemotePtrType, r *sstore.RemoteType) *WebShareRemote { + return &WebShareRemote{ + RemoteId: r.RemoteId, + Alias: r.RemoteAlias, + CanonicalName: r.RemoteCanonicalName, + Name: rptr.Name, + HomeDir: r.StateVars["home"], + IsRoot: r.StateVars["remoteuser"] == "root", + } +} + +func webScreenFromScreen(s *sstore.ScreenType) (*WebShareScreenType, error) { + if s == nil || s.ScreenId == "" { + return nil, fmt.Errorf("invalid nil screen") + } + if s.WebShareOpts == nil { + return nil, fmt.Errorf("invalid screen, no WebShareOpts") + } + if s.WebShareOpts.ViewKey == "" { + return nil, fmt.Errorf("invalid screen, no ViewKey") + } + var shareName string + if s.WebShareOpts.ShareName != "" { + shareName = s.WebShareOpts.ShareName + } else { + shareName = s.Name + } + return &WebShareScreenType{ScreenId: s.ScreenId, ShareName: shareName, ViewKey: s.WebShareOpts.ViewKey, SelectedLine: int(s.SelectedLine)}, nil +} + +type WebShareLineType struct { + LineId string `json:"lineid"` + Ts int64 `json:"ts"` + LineNum int64 `json:"linenum"` + LineType string `json:"linetype"` + ContentHeight int64 `json:"contentheight"` + Renderer string `json:"renderer,omitempty"` + Text string `json:"text,omitempty"` +} + +func webLineFromLine(line *sstore.LineType) (*WebShareLineType, error) { + rtn := &WebShareLineType{ + LineId: line.LineId, + Ts: line.Ts, + LineNum: line.LineNum, + LineType: line.LineType, + ContentHeight: line.ContentHeight, + Renderer: line.Renderer, + Text: line.Text, + } + return rtn, nil +} + +type WebShareCmdType struct { + LineId string `json:"lineid"` + CmdStr string `json:"cmdstr"` + RawCmdStr string `json:"rawcmdstr"` + Remote *WebShareRemote `json:"remote"` + FeState sstore.FeStateType `json:"festate"` + TermOpts sstore.TermOpts `json:"termopts"` + Status string `json:"status"` + CmdPid int `json:"cmdpid"` + RemotePid int `json:"remotepid"` + DoneTs int64 `json:"donets,omitempty"` + ExitCode int `json:"exitcode,omitempty"` + DurationMs int `json:"durationms,omitempty"` + RtnState bool `json:"rtnstate,omitempty"` + RtnStateStr string `json:"rtnstatestr,omitempty"` +} + +func webCmdFromCmd(lineId string, cmd *sstore.CmdType) (*WebShareCmdType, error) { + if cmd.Remote.RemoteId == "" { + return nil, fmt.Errorf("invalid cmd, remoteptr has no remoteid") + } + remote := remote.GetRemoteCopyById(cmd.Remote.RemoteId) + if remote == nil { + return nil, fmt.Errorf("invalid cmd, cannot retrieve remote:%s", cmd.Remote.RemoteId) + } + webRemote := webRemoteFromRemote(cmd.Remote, remote) + rtn := &WebShareCmdType{ + LineId: lineId, + CmdStr: cmd.CmdStr, + RawCmdStr: cmd.RawCmdStr, + Remote: webRemote, + FeState: cmd.FeState, + TermOpts: cmd.TermOpts, + Status: cmd.Status, + CmdPid: cmd.CmdPid, + RemotePid: cmd.RemotePid, + DoneTs: cmd.DoneTs, + ExitCode: cmd.ExitCode, + DurationMs: cmd.DurationMs, + RtnState: cmd.RtnState, + } + if cmd.RtnState { + barr, err := rtnstate.GetRtnStateDiff(context.Background(), cmd.ScreenId, cmd.LineId) + if err != nil { + return nil, fmt.Errorf("error creating rtnstate diff for cmd:%s: %v", cmd.LineId, err) + } + rtn.RtnStateStr = string(barr) + } + return rtn, nil +} + +type WebSharePtyData struct { + PtyPos int64 `json:"ptypos"` + Data []byte `json:"data"` + Eof bool `json:"-"` // internal use +} diff --git a/wavesrv/pkg/promptenc/promptenc.go b/wavesrv/pkg/promptenc/promptenc.go new file mode 100644 index 00000000..e5ac25b3 --- /dev/null +++ b/wavesrv/pkg/promptenc/promptenc.go @@ -0,0 +1,199 @@ +package promptenc + +import ( + "crypto/cipher" + "crypto/rand" + "encoding/base64" + "encoding/json" + "fmt" + "io" + "reflect" + + ccp "golang.org/x/crypto/chacha20poly1305" +) + +const EncTagName = "enc" +const EncFieldIndicator = "*" + +type Encryptor struct { + Key []byte + AEAD cipher.AEAD +} + +type HasOData interface { + GetOData() string +} + +func readRandBytes(n int) ([]byte, error) { + rtn := make([]byte, n) + _, err := io.ReadFull(rand.Reader, rtn) + return rtn, err +} + +func MakeRandomEncryptor() (*Encryptor, error) { + key, err := readRandBytes(ccp.KeySize) + if err != nil { + return nil, err + } + rtn := &Encryptor{Key: key} + rtn.AEAD, err = ccp.NewX(rtn.Key) + if err != nil { + return nil, err + } + return rtn, nil +} + +func MakeEncryptor(key []byte) (*Encryptor, error) { + var err error + rtn := &Encryptor{Key: key} + rtn.AEAD, err = ccp.NewX(rtn.Key) + if err != nil { + return nil, err + } + return rtn, nil +} + +func MakeEncryptorB64(key64 string) (*Encryptor, error) { + keyBytes, err := base64.RawURLEncoding.DecodeString(key64) + if err != nil { + return nil, err + } + return MakeEncryptor(keyBytes) +} + +func (enc *Encryptor) EncryptData(plainText []byte, odata string) ([]byte, error) { + outputBuf := make([]byte, enc.AEAD.NonceSize()+enc.AEAD.Overhead()+len(plainText)) + nonce := outputBuf[0:enc.AEAD.NonceSize()] + _, err := io.ReadFull(rand.Reader, nonce) + if err != nil { + return nil, err + } + // we're going to append the cipherText to nonce. so the encrypted data is [nonce][ciphertext] + // note that outputbuf should be the correct size to hold the rtn value + rtn := enc.AEAD.Seal(nonce, nonce, plainText, []byte(odata)) + return rtn, nil +} + +func (enc *Encryptor) DecryptData(encData []byte, odata string) (map[string]interface{}, error) { + minLen := enc.AEAD.NonceSize() + enc.AEAD.Overhead() + if len(encData) < minLen { + return nil, fmt.Errorf("invalid encdata, len:%d is less than minimum len:%d", len(encData), minLen) + } + m := make(map[string]interface{}) + nonce := encData[0:enc.AEAD.NonceSize()] + cipherText := encData[enc.AEAD.NonceSize():] + plainText, err := enc.AEAD.Open(nil, nonce, cipherText, []byte(odata)) + if err != nil { + return nil, err + } + err = json.Unmarshal(plainText, &m) + if err != nil { + return nil, err + } + return m, nil +} + +type EncryptMeta struct { + EncField *reflect.StructField + PlainFields map[string]reflect.StructField +} + +func isByteArrayType(t reflect.Type) bool { + return t.Kind() == reflect.Slice && t.Elem().Kind() == reflect.Uint8 +} + +func metaFromType(v interface{}) (*EncryptMeta, error) { + if v == nil { + return nil, fmt.Errorf("Encryptor cannot encrypt nil") + } + rt := reflect.TypeOf(v) + if rt.Kind() != reflect.Pointer { + return nil, fmt.Errorf("Encryptor invalid type %T, not a pointer type", v) + } + rtElem := rt.Elem() + if rtElem.Kind() != reflect.Struct { + return nil, fmt.Errorf("Encryptor invalid type %T, not a pointer to struct type", v) + } + meta := &EncryptMeta{} + meta.PlainFields = make(map[string]reflect.StructField) + numFields := rtElem.NumField() + for i := 0; i < numFields; i++ { + field := rtElem.Field(i) + encTag := field.Tag.Get(EncTagName) + if encTag == "" { + continue + } + if encTag == EncFieldIndicator { + if meta.EncField != nil { + return nil, fmt.Errorf("Encryptor, type %T has two enc fields set (*)", v) + } + if !isByteArrayType(field.Type) { + return nil, fmt.Errorf("Encryptor, type %T enc field %q is not []byte", v, field.Name) + } + meta.EncField = &field + continue + } + if _, found := meta.PlainFields[encTag]; found { + return nil, fmt.Errorf("Encryptor, type %T has two enc fields with tag %q", v, encTag) + } + meta.PlainFields[encTag] = field + } + if meta.EncField == nil { + return nil, fmt.Errorf("Encryptor, type %T has no enc (*) field", v) + } + return meta, nil +} + +func (enc *Encryptor) EncryptODS(v HasOData) error { + odata := v.GetOData() + return enc.EncryptStructFields(v, odata) +} + +func (enc *Encryptor) DecryptODS(v HasOData) error { + odata := v.GetOData() + return enc.DecryptStructFields(v, odata) +} + +func (enc *Encryptor) EncryptStructFields(v interface{}, odata string) error { + encMeta, err := metaFromType(v) + if err != nil { + return err + } + rvPtr := reflect.ValueOf(v) + rv := rvPtr.Elem() + m := make(map[string]interface{}) + for jsonKey, field := range encMeta.PlainFields { + fieldVal := rv.FieldByIndex(field.Index) + m[jsonKey] = fieldVal.Interface() + } + barr, err := json.Marshal(m) + if err != nil { + return err + } + cipherText, err := enc.EncryptData(barr, odata) + if err != nil { + return err + } + encFieldValue := rv.FieldByIndex(encMeta.EncField.Index) + encFieldValue.SetBytes(cipherText) + return nil +} + +func (enc *Encryptor) DecryptStructFields(v interface{}, odata string) error { + encMeta, err := metaFromType(v) + if err != nil { + return err + } + rvPtr := reflect.ValueOf(v) + rv := rvPtr.Elem() + cipherText := rv.FieldByIndex(encMeta.EncField.Index).Bytes() + m, err := enc.DecryptData(cipherText, odata) + if err != nil { + return err + } + for jsonKey, field := range encMeta.PlainFields { + val := m[jsonKey] + rv.FieldByIndex(field.Index).Set(reflect.ValueOf(val)) + } + return nil +} diff --git a/wavesrv/pkg/remote/circlelog.go b/wavesrv/pkg/remote/circlelog.go new file mode 100644 index 00000000..b5d3eb29 --- /dev/null +++ b/wavesrv/pkg/remote/circlelog.go @@ -0,0 +1,53 @@ +package remote + +import ( + "fmt" + "sync" +) + +type CircleLog struct { + Lock *sync.Mutex + StartPos int + Log []string + MaxSize int +} + +func MakeCircleLog(maxSize int) *CircleLog { + if maxSize <= 0 { + panic("invalid maxsize, must be >= 0") + } + rtn := &CircleLog{ + Lock: &sync.Mutex{}, + StartPos: 0, + Log: make([]string, 0, maxSize), + MaxSize: maxSize, + } + return rtn +} + +func (l *CircleLog) Add(s string) { + l.Lock.Lock() + defer l.Lock.Unlock() + if len(l.Log) < l.MaxSize { + l.Log = append(l.Log, s) + return + } + l.Log[l.StartPos] = s + l.StartPos = (l.StartPos + 1) % l.MaxSize +} + +func (l *CircleLog) Addf(sfmt string, args ...interface{}) { + // no lock here, since l.Add() is synchronized + s := fmt.Sprintf(sfmt, args...) + l.Add(s) +} + +func (l *CircleLog) GetEntries() []string { + l.Lock.Lock() + defer l.Lock.Unlock() + rtn := make([]string, len(l.Log)) + for i := 0; i < len(l.Log); i++ { + rtn[i] = l.Log[(l.StartPos+i)%l.MaxSize] + } + return rtn +} diff --git a/wavesrv/pkg/remote/openai/openai.go b/wavesrv/pkg/remote/openai/openai.go new file mode 100644 index 00000000..a2f646b5 --- /dev/null +++ b/wavesrv/pkg/remote/openai/openai.go @@ -0,0 +1,147 @@ +package openai + +import ( + "context" + "fmt" + "io" + + openaiapi "github.com/sashabaranov/go-openai" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/prompt-server/pkg/sstore" +) + +// https://github.com/tiktoken-go/tokenizer + +const DefaultMaxTokens = 1000 +const DefaultModel = "gpt-3.5-turbo" +const DefaultStreamChanSize = 10 + +func convertUsage(resp openaiapi.ChatCompletionResponse) *packet.OpenAIUsageType { + if resp.Usage.TotalTokens == 0 { + return nil + } + return &packet.OpenAIUsageType{ + PromptTokens: resp.Usage.PromptTokens, + CompletionTokens: resp.Usage.CompletionTokens, + TotalTokens: resp.Usage.TotalTokens, + } +} + +func convertPrompt(prompt []sstore.OpenAIPromptMessageType) []openaiapi.ChatCompletionMessage { + var rtn []openaiapi.ChatCompletionMessage + for _, p := range prompt { + msg := openaiapi.ChatCompletionMessage{Role: p.Role, Content: p.Content, Name: p.Name} + rtn = append(rtn, msg) + } + return rtn +} + +func RunCompletion(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) ([]*packet.OpenAIPacketType, error) { + if opts == nil { + return nil, fmt.Errorf("no openai opts found") + } + if opts.Model == "" { + return nil, fmt.Errorf("no openai model specified") + } + if opts.APIToken == "" { + return nil, fmt.Errorf("no api token") + } + client := openaiapi.NewClient(opts.APIToken) + req := openaiapi.ChatCompletionRequest{ + Model: opts.Model, + Messages: convertPrompt(prompt), + MaxTokens: opts.MaxTokens, + } + if opts.MaxChoices > 1 { + req.N = opts.MaxChoices + } + apiResp, err := client.CreateChatCompletion(ctx, req) + if err != nil { + return nil, fmt.Errorf("error calling openai API: %v", err) + } + if len(apiResp.Choices) == 0 { + return nil, fmt.Errorf("no response received") + } + return marshalResponse(apiResp), nil +} + +func RunCompletionStream(ctx context.Context, opts *sstore.OpenAIOptsType, prompt []sstore.OpenAIPromptMessageType) (chan *packet.OpenAIPacketType, error) { + if opts == nil { + return nil, fmt.Errorf("no openai opts found") + } + if opts.Model == "" { + return nil, fmt.Errorf("no openai model specified") + } + if opts.APIToken == "" { + return nil, fmt.Errorf("no api token") + } + client := openaiapi.NewClient(opts.APIToken) + req := openaiapi.ChatCompletionRequest{ + Model: opts.Model, + Messages: convertPrompt(prompt), + MaxTokens: opts.MaxTokens, + Stream: true, + } + if opts.MaxChoices > 1 { + req.N = opts.MaxChoices + } + apiResp, err := client.CreateChatCompletionStream(ctx, req) + if err != nil { + return nil, fmt.Errorf("error calling openai API: %v", err) + } + rtn := make(chan *packet.OpenAIPacketType, DefaultStreamChanSize) + go func() { + sentHeader := false + defer close(rtn) + for { + streamResp, err := apiResp.Recv() + if err == io.EOF { + break + } + if err != nil { + errPk := CreateErrorPacket(fmt.Sprintf("error in recv of streaming data: %v", err)) + rtn <- errPk + break + } + if streamResp.Model != "" && !sentHeader { + pk := packet.MakeOpenAIPacket() + pk.Model = streamResp.Model + pk.Created = streamResp.Created + rtn <- pk + sentHeader = true + } + for _, choice := range streamResp.Choices { + pk := packet.MakeOpenAIPacket() + pk.Index = choice.Index + pk.Text = choice.Delta.Content + pk.FinishReason = choice.FinishReason + rtn <- pk + } + } + }() + return rtn, err +} + +func marshalResponse(resp openaiapi.ChatCompletionResponse) []*packet.OpenAIPacketType { + var rtn []*packet.OpenAIPacketType + headerPk := packet.MakeOpenAIPacket() + headerPk.Model = resp.Model + headerPk.Created = resp.Created + headerPk.Usage = convertUsage(resp) + rtn = append(rtn, headerPk) + for _, choice := range resp.Choices { + choicePk := packet.MakeOpenAIPacket() + choicePk.Index = choice.Index + choicePk.Text = choice.Message.Content + choicePk.FinishReason = choice.FinishReason + rtn = append(rtn, choicePk) + } + return rtn +} + +func CreateErrorPacket(errStr string) *packet.OpenAIPacketType { + errPk := packet.MakeOpenAIPacket() + errPk.FinishReason = "error" + errPk.Error = errStr + return errPk +} diff --git a/wavesrv/pkg/remote/remote.go b/wavesrv/pkg/remote/remote.go new file mode 100644 index 00000000..5fcdb409 --- /dev/null +++ b/wavesrv/pkg/remote/remote.go @@ -0,0 +1,2104 @@ +package remote + +import ( + "bytes" + "context" + "encoding/base64" + "errors" + "fmt" + "io" + "log" + "os" + "os/exec" + "path" + "regexp" + "strconv" + "strings" + "sync" + "syscall" + "time" + + "github.com/armon/circbuf" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/apishell/pkg/statediff" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/commandlinedev/prompt-server/pkg/scpacket" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/creack/pty" + "github.com/google/uuid" + "golang.org/x/mod/semver" +) + +const RemoteTypeMShell = "mshell" +const DefaultTerm = "xterm-256color" +const DefaultMaxPtySize = 1024 * 1024 +const CircBufSize = 64 * 1024 +const RemoteTermRows = 8 +const RemoteTermCols = 80 +const PtyReadBufSize = 100 +const RemoteConnectTimeout = 15 * time.Second + +const MShellServerCommandFmt = ` +PATH=$PATH:~/.mshell; +which mshell-[%VERSION%] > /dev/null; +if [[ "$?" -ne 0 ]] +then + printf "\n##N{\"type\": \"init\", \"notfound\": true, \"uname\": \"%s | %s\"}\n" "$(uname -s)" "$(uname -m)" +else + mshell-[%VERSION%] --server +fi +` + +func MakeLocalMShellCommandStr(isSudo bool) (string, error) { + mshellPath, err := scbase.LocalMShellBinaryPath() + if err != nil { + return "", err + } + if isSudo { + return fmt.Sprintf("sudo %s --server", mshellPath), nil + } else { + return fmt.Sprintf("%s --server", mshellPath), nil + } +} + +func MakeServerCommandStr() string { + return strings.ReplaceAll(MShellServerCommandFmt, "[%VERSION%]", semver.MajorMinor(scbase.MShellVersion)) +} + +const ( + StatusConnected = "connected" + StatusConnecting = "connecting" + StatusDisconnected = "disconnected" + StatusError = "error" +) + +func init() { + if scbase.MShellVersion != base.MShellVersion { + panic(fmt.Sprintf("prompt-server apishell version must match '%s' vs '%s'", scbase.MShellVersion, base.MShellVersion)) + } +} + +var GlobalStore *Store + +type Store struct { + Lock *sync.Mutex + Map map[string]*MShellProc // key=remoteid + CmdWaitMap map[base.CommandKey][]func() +} + +type MShellProc struct { + Lock *sync.Mutex + Remote *sstore.RemoteType + + // runtime + RemoteId string // can be read without a lock + Status string + ServerProc *shexec.ClientProc + UName string + Err error + ErrNoInitPk bool + ControllingPty *os.File + PtyBuffer *circbuf.Buffer + MakeClientCancelFn context.CancelFunc + MakeClientDeadline *time.Time + StateMap map[string]*packet.ShellState // sha1->state + CurrentState string // sha1 + NumTryConnect int + + // install + InstallStatus string + NeedsMShellUpgrade bool + InstallCancelFn context.CancelFunc + InstallErr error + + RunningCmds map[base.CommandKey]RunCmdType + WaitingCmds []RunCmdType + PendingStateCmds map[string]base.CommandKey // key=[remoteinstance name] +} + +type RunCmdType struct { + SessionId string + ScreenId string + RemotePtr sstore.RemotePtrType + RunPacket *packet.RunPacketType +} + +type RemoteRuntimeState struct { + RemoteType string `json:"remotetype"` + RemoteId string `json:"remoteid"` + RemoteAlias string `json:"remotealias,omitempty"` + RemoteCanonicalName string `json:"remotecanonicalname"` + RemoteVars map[string]string `json:"remotevars"` + DefaultFeState map[string]string `json:"defaultfestate"` + Status string `json:"status"` + ConnectTimeout int `json:"connecttimeout,omitempty"` + ErrorStr string `json:"errorstr,omitempty"` + InstallStatus string `json:"installstatus"` + InstallErrorStr string `json:"installerrorstr,omitempty"` + NeedsMShellUpgrade bool `json:"needsmshellupgrade,omitempty"` + NoInitPk bool `json:"noinitpk,omitempty"` + AuthType string `json:"authtype,omitempty"` + ConnectMode string `json:"connectmode"` + AutoInstall bool `json:"autoinstall"` + Archived bool `json:"archived,omitempty"` + RemoteIdx int64 `json:"remoteidx"` + UName string `json:"uname"` + MShellVersion string `json:"mshellversion"` + WaitingForPassword bool `json:"waitingforpassword,omitempty"` + Local bool `json:"local,omitempty"` + RemoteOpts *sstore.RemoteOptsType `json:"remoteopts,omitempty"` + CanComplete bool `json:"cancomplete,omitempty"` +} + +func (state RemoteRuntimeState) IsConnected() bool { + return state.Status == StatusConnected +} + +func CanComplete(remoteType string) bool { + switch remoteType { + case sstore.RemoteTypeSsh: + return true + default: + return false + } +} + +func (msh *MShellProc) GetStatus() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Status +} + +func (msh *MShellProc) GetDefaultState() *packet.ShellState { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.StateMap[msh.CurrentState] +} + +func (msh *MShellProc) GetDefaultStatePtr() *sstore.ShellStatePtr { + msh.Lock.Lock() + defer msh.Lock.Unlock() + if msh.CurrentState == "" { + return nil + } + return &sstore.ShellStatePtr{BaseHash: msh.CurrentState} +} + +func (msh *MShellProc) GetDefaultFeState() map[string]string { + state := msh.GetDefaultState() + return sstore.FeStateFromShellState(state) +} + +func (msh *MShellProc) GetStateByHash(hval string) *packet.ShellState { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.StateMap[hval] +} + +func (msh *MShellProc) GetRemoteId() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Remote.RemoteId +} + +func (msh *MShellProc) GetInstallStatus() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.InstallStatus +} + +func (state RemoteRuntimeState) GetBaseDisplayName() string { + if state.RemoteAlias != "" { + return state.RemoteAlias + } + return state.RemoteCanonicalName +} + +func (state RemoteRuntimeState) GetDisplayName(rptr *sstore.RemotePtrType) string { + baseDisplayName := state.GetBaseDisplayName() + if rptr == nil { + return baseDisplayName + } + return rptr.GetDisplayName(baseDisplayName) +} + +func LoadRemotes(ctx context.Context) error { + GlobalStore = &Store{ + Lock: &sync.Mutex{}, + Map: make(map[string]*MShellProc), + CmdWaitMap: make(map[base.CommandKey][]func()), + } + allRemotes, err := sstore.GetAllRemotes(ctx) + if err != nil { + return err + } + var numLocal int + var numSudoLocal int + for _, remote := range allRemotes { + msh := MakeMShell(remote) + GlobalStore.Map[remote.RemoteId] = msh + if remote.ConnectMode == sstore.ConnectModeStartup { + go msh.Launch(false) + } + if remote.Local { + if remote.IsSudo() { + numSudoLocal++ + } else { + numLocal++ + } + } + } + if numLocal == 0 { + return fmt.Errorf("no local remote found") + } + if numLocal > 1 { + return fmt.Errorf("multiple local remotes found") + } + if numSudoLocal > 1 { + return fmt.Errorf("multiple local sudo remotes found") + } + return nil +} + +func LoadRemoteById(ctx context.Context, remoteId string) error { + r, err := sstore.GetRemoteById(ctx, remoteId) + if err != nil { + return err + } + if r == nil { + return fmt.Errorf("remote %s not found", remoteId) + } + msh := MakeMShell(r) + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + existingRemote := GlobalStore.Map[remoteId] + if existingRemote != nil { + return fmt.Errorf("cannot add remote %s, already in global map", remoteId) + } + GlobalStore.Map[r.RemoteId] = msh + if r.ConnectMode == sstore.ConnectModeStartup { + go msh.Launch(false) + } + return nil +} + +func ReadRemotePty(ctx context.Context, remoteId string) (int64, []byte, error) { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + msh := GlobalStore.Map[remoteId] + if msh == nil { + return 0, nil, nil + } + msh.Lock.Lock() + defer msh.Lock.Unlock() + barr := msh.PtyBuffer.Bytes() + offset := msh.PtyBuffer.TotalWritten() - int64(len(barr)) + return offset, barr, nil +} + +func AddRemote(ctx context.Context, r *sstore.RemoteType, shouldStart bool) error { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + + existingRemote := getRemoteByCanonicalName_nolock(r.RemoteCanonicalName) + if existingRemote != nil { + erCopy := existingRemote.GetRemoteCopy() + if !erCopy.Archived { + return fmt.Errorf("duplicate canonical name %q: cannot create new remote", r.RemoteCanonicalName) + } + r.RemoteId = erCopy.RemoteId + } + if r.Local { + return fmt.Errorf("cannot create another local remote (there can be only one)") + } + + err := sstore.UpsertRemote(ctx, r) + if err != nil { + return fmt.Errorf("cannot create remote %q: %v", r.RemoteCanonicalName, err) + } + newMsh := MakeMShell(r) + GlobalStore.Map[r.RemoteId] = newMsh + go newMsh.NotifyRemoteUpdate() + if shouldStart { + go newMsh.Launch(true) + } + return nil +} + +func ArchiveRemote(ctx context.Context, remoteId string) error { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + msh := GlobalStore.Map[remoteId] + if msh == nil { + return fmt.Errorf("remote not found, cannot archive") + } + if msh.Status == StatusConnected { + return fmt.Errorf("cannot archive connected remote") + } + if msh.Remote.Local { + return fmt.Errorf("cannot archive local remote") + } + rcopy := msh.GetRemoteCopy() + archivedRemote := &sstore.RemoteType{ + RemoteId: rcopy.RemoteId, + RemoteType: rcopy.RemoteType, + RemoteCanonicalName: rcopy.RemoteCanonicalName, + ConnectMode: sstore.ConnectModeManual, + Archived: true, + } + err := sstore.UpsertRemote(ctx, archivedRemote) + if err != nil { + return err + } + newMsh := MakeMShell(archivedRemote) + GlobalStore.Map[remoteId] = newMsh + go newMsh.NotifyRemoteUpdate() + return nil +} + +var partialUUIDRe = regexp.MustCompile("^[0-9a-f]{8}$") + +func isPartialUUID(s string) bool { + return partialUUIDRe.MatchString(s) +} + +func NumRemotes() int { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + return len(GlobalStore.Map) +} + +func GetRemoteByArg(arg string) *MShellProc { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + isPuid := isPartialUUID(arg) + for _, msh := range GlobalStore.Map { + rcopy := msh.GetRemoteCopy() + if rcopy.RemoteAlias == arg || rcopy.RemoteCanonicalName == arg || rcopy.RemoteId == arg { + return msh + } + if isPuid && strings.HasPrefix(rcopy.RemoteId, arg) { + return msh + } + } + return nil +} + +func getRemoteByCanonicalName_nolock(name string) *MShellProc { + for _, msh := range GlobalStore.Map { + rcopy := msh.GetRemoteCopy() + if rcopy.RemoteCanonicalName == name { + return msh + } + } + return nil +} + +func GetRemoteById(remoteId string) *MShellProc { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + return GlobalStore.Map[remoteId] +} + +func GetRemoteCopyById(remoteId string) *sstore.RemoteType { + msh := GetRemoteById(remoteId) + if msh == nil { + return nil + } + rcopy := msh.GetRemoteCopy() + return &rcopy +} + +func GetRemoteMap() map[string]*MShellProc { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + rtn := make(map[string]*MShellProc) + for remoteId, msh := range GlobalStore.Map { + rtn[remoteId] = msh + } + return rtn +} + +func GetLocalRemote() *MShellProc { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + for _, msh := range GlobalStore.Map { + if msh.IsLocal() && !msh.IsSudo() { + return msh + } + } + return nil +} + +func ResolveRemoteRef(remoteRef string) *RemoteRuntimeState { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + + _, err := uuid.Parse(remoteRef) + if err == nil { + msh := GlobalStore.Map[remoteRef] + if msh != nil { + state := msh.GetRemoteRuntimeState() + return &state + } + return nil + } + for _, msh := range GlobalStore.Map { + if msh.Remote.RemoteAlias == remoteRef || msh.Remote.RemoteCanonicalName == remoteRef { + state := msh.GetRemoteRuntimeState() + return &state + } + } + return nil +} + +func unquoteDQBashString(str string) (string, bool) { + if len(str) < 2 { + return str, false + } + if str[0] != '"' || str[len(str)-1] != '"' { + return str, false + } + rtn := make([]byte, 0, len(str)-2) + for idx := 1; idx < len(str)-1; idx++ { + ch := str[idx] + if ch == '"' { + return str, false + } + if ch == '\\' { + if idx == len(str)-2 { + return str, false + } + nextCh := str[idx+1] + if nextCh == '\n' { + idx++ + continue + } + if nextCh == '$' || nextCh == '"' || nextCh == '\\' || nextCh == '`' { + idx++ + rtn = append(rtn, nextCh) + continue + } + rtn = append(rtn, '\\') + continue + } else { + rtn = append(rtn, ch) + } + } + return string(rtn), true +} + +func makeShortHost(host string) string { + dotIdx := strings.Index(host, ".") + if dotIdx == -1 { + return host + } + return host[0:dotIdx] +} + +func (msh *MShellProc) IsLocal() bool { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Remote.Local +} + +func (msh *MShellProc) IsSudo() bool { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Remote.IsSudo() +} + +func (msh *MShellProc) tryAutoInstall() { + msh.Lock.Lock() + defer msh.Lock.Unlock() + if !msh.Remote.AutoInstall || !msh.NeedsMShellUpgrade || msh.InstallErr != nil { + return + } + msh.writeToPtyBuffer_nolock("trying auto-install\n") + go msh.RunInstall() +} + +func (msh *MShellProc) GetRemoteRuntimeState() RemoteRuntimeState { + msh.Lock.Lock() + defer msh.Lock.Unlock() + state := RemoteRuntimeState{ + RemoteType: msh.Remote.RemoteType, + RemoteId: msh.Remote.RemoteId, + RemoteAlias: msh.Remote.RemoteAlias, + RemoteCanonicalName: msh.Remote.RemoteCanonicalName, + Status: msh.Status, + ConnectMode: msh.Remote.ConnectMode, + AutoInstall: msh.Remote.AutoInstall, + Archived: msh.Remote.Archived, + RemoteIdx: msh.Remote.RemoteIdx, + UName: msh.UName, + InstallStatus: msh.InstallStatus, + NeedsMShellUpgrade: msh.NeedsMShellUpgrade, + Local: msh.Remote.Local, + NoInitPk: msh.ErrNoInitPk, + AuthType: sstore.RemoteAuthTypeNone, + } + if msh.Remote.SSHOpts != nil { + state.AuthType = msh.Remote.SSHOpts.GetAuthType() + } + if msh.Remote.RemoteOpts != nil { + optsCopy := *msh.Remote.RemoteOpts + state.RemoteOpts = &optsCopy + } + if msh.Err != nil { + state.ErrorStr = msh.Err.Error() + } + if msh.InstallErr != nil { + state.InstallErrorStr = msh.InstallErr.Error() + } + if msh.Status == StatusConnecting { + state.WaitingForPassword = msh.isWaitingForPassword_nolock() + if msh.MakeClientDeadline != nil { + state.ConnectTimeout = int((*msh.MakeClientDeadline).Sub(time.Now()) / time.Second) + if state.ConnectTimeout < 0 { + state.ConnectTimeout = 0 + } + } + } + vars := msh.Remote.StateVars + if vars == nil { + vars = make(map[string]string) + } + vars["user"] = msh.Remote.RemoteUser + vars["bestuser"] = vars["user"] + vars["host"] = msh.Remote.RemoteHost + vars["shorthost"] = makeShortHost(msh.Remote.RemoteHost) + vars["alias"] = msh.Remote.RemoteAlias + vars["cname"] = msh.Remote.RemoteCanonicalName + vars["remoteid"] = msh.Remote.RemoteId + vars["status"] = msh.Status + vars["type"] = msh.Remote.RemoteType + if msh.Remote.IsSudo() { + vars["sudo"] = "1" + } + if msh.Remote.Local { + vars["local"] = "1" + } + vars["port"] = "22" + if msh.Remote.SSHOpts != nil { + if msh.Remote.SSHOpts.SSHPort != 0 { + vars["port"] = strconv.Itoa(msh.Remote.SSHOpts.SSHPort) + } + } + if msh.Remote.RemoteOpts != nil && msh.Remote.RemoteOpts.Color != "" { + vars["color"] = msh.Remote.RemoteOpts.Color + } + if msh.ServerProc != nil && msh.ServerProc.InitPk != nil { + initPk := msh.ServerProc.InitPk + if initPk.BuildTime == "" || initPk.BuildTime == "0" { + state.MShellVersion = initPk.Version + } else { + state.MShellVersion = fmt.Sprintf("%s+%s", initPk.Version, initPk.BuildTime) + } + vars["home"] = initPk.HomeDir + vars["remoteuser"] = initPk.User + vars["bestuser"] = vars["remoteuser"] + vars["remotehost"] = initPk.HostName + vars["remoteshorthost"] = makeShortHost(initPk.HostName) + vars["besthost"] = vars["remotehost"] + vars["bestshorthost"] = vars["remoteshorthost"] + } + curState := msh.StateMap[msh.CurrentState] + if curState != nil { + state.DefaultFeState = sstore.FeStateFromShellState(curState) + vars["cwd"] = curState.Cwd + } + if msh.Remote.Local && msh.Remote.IsSudo() { + vars["bestuser"] = "sudo" + } else if msh.Remote.IsSudo() { + vars["bestuser"] = "sudo@" + vars["bestuser"] + } + if msh.Remote.Local { + vars["bestname"] = vars["bestuser"] + "@local" + vars["bestshortname"] = vars["bestuser"] + "@local" + } else { + vars["bestname"] = vars["bestuser"] + "@" + vars["besthost"] + vars["bestshortname"] = vars["bestuser"] + "@" + vars["bestshorthost"] + } + if vars["remoteuser"] == "root" || vars["sudo"] == "1" { + vars["isroot"] = "1" + } + state.RemoteVars = vars + return state +} + +func (msh *MShellProc) NotifyRemoteUpdate() { + rstate := msh.GetRemoteRuntimeState() + update := &sstore.ModelUpdate{Remotes: []interface{}{rstate}} + sstore.MainBus.SendUpdate(update) +} + +func GetAllRemoteRuntimeState() []RemoteRuntimeState { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + + var rtn []RemoteRuntimeState + for _, proc := range GlobalStore.Map { + state := proc.GetRemoteRuntimeState() + rtn = append(rtn, state) + } + return rtn +} + +func GetDefaultRemoteStateById(remoteId string) (*packet.ShellState, error) { + remote := GetRemoteById(remoteId) + if remote == nil { + return nil, fmt.Errorf("remote not found") + } + if !remote.IsConnected() { + return nil, fmt.Errorf("remote not connected") + } + state := remote.GetDefaultState() + if state == nil { + return nil, fmt.Errorf("could not get default remote state") + } + return state, nil +} + +func MakeMShell(r *sstore.RemoteType) *MShellProc { + buf, err := circbuf.NewBuffer(CircBufSize) + if err != nil { + panic(err) // this should never happen (NewBuffer only returns an error if CirBufSize <= 0) + } + rtn := &MShellProc{ + Lock: &sync.Mutex{}, + Remote: r, + RemoteId: r.RemoteId, + Status: StatusDisconnected, + PtyBuffer: buf, + InstallStatus: StatusDisconnected, + RunningCmds: make(map[base.CommandKey]RunCmdType), + PendingStateCmds: make(map[string]base.CommandKey), + StateMap: make(map[string]*packet.ShellState), + } + rtn.WriteToPtyBuffer("console for connection [%s]\n", r.GetName()) + return rtn +} + +func SendRemoteInput(pk *scpacket.RemoteInputPacketType) error { + data, err := base64.StdEncoding.DecodeString(pk.InputData64) + if err != nil { + return fmt.Errorf("cannot decode base64: %v\n", err) + } + msh := GetRemoteById(pk.RemoteId) + if msh == nil { + return fmt.Errorf("remote not found") + } + var cmdPty *os.File + msh.WithLock(func() { + cmdPty = msh.ControllingPty + }) + if cmdPty == nil { + return fmt.Errorf("remote has no attached pty") + } + _, err = cmdPty.Write(data) + if err != nil { + return fmt.Errorf("writing to pty: %v", err) + } + msh.resetClientDeadline() + return nil +} + +func (msh *MShellProc) getClientDeadline() *time.Time { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.MakeClientDeadline +} + +func (msh *MShellProc) resetClientDeadline() { + msh.Lock.Lock() + defer msh.Lock.Unlock() + if msh.Status != StatusConnecting { + return + } + deadline := msh.MakeClientDeadline + if deadline == nil { + return + } + newDeadline := time.Now().Add(RemoteConnectTimeout) + msh.MakeClientDeadline = &newDeadline +} + +func (msh *MShellProc) watchClientDeadlineTime() { + for { + time.Sleep(1 * time.Second) + status := msh.GetStatus() + if status != StatusConnecting { + break + } + deadline := msh.getClientDeadline() + if deadline == nil { + break + } + if time.Now().After(*deadline) { + msh.Disconnect(false) + break + } + go msh.NotifyRemoteUpdate() + } +} + +func convertSSHOpts(opts *sstore.SSHOpts) shexec.SSHOpts { + if opts == nil || opts.Local { + opts = &sstore.SSHOpts{} + } + return shexec.SSHOpts{ + SSHHost: opts.SSHHost, + SSHOptsStr: opts.SSHOptsStr, + SSHIdentity: opts.SSHIdentity, + SSHUser: opts.SSHUser, + SSHPort: opts.SSHPort, + } +} + +func (msh *MShellProc) addControllingTty(ecmd *exec.Cmd) (*os.File, error) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + + cmdPty, cmdTty, err := pty.Open() + if err != nil { + return nil, err + } + pty.Setsize(cmdPty, &pty.Winsize{Rows: RemoteTermRows, Cols: RemoteTermCols}) + msh.ControllingPty = cmdPty + ecmd.ExtraFiles = append(ecmd.ExtraFiles, cmdTty) + if ecmd.SysProcAttr == nil { + ecmd.SysProcAttr = &syscall.SysProcAttr{} + } + ecmd.SysProcAttr.Setsid = true + ecmd.SysProcAttr.Setctty = true + ecmd.SysProcAttr.Ctty = len(ecmd.ExtraFiles) + 3 - 1 + return cmdPty, nil +} + +func (msh *MShellProc) setErrorStatus(err error) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + msh.Status = StatusError + msh.Err = err + go msh.NotifyRemoteUpdate() +} + +func (msh *MShellProc) setInstallErrorStatus(err error) { + msh.WriteToPtyBuffer("*error, %s\n", err.Error()) + msh.Lock.Lock() + defer msh.Lock.Unlock() + msh.InstallStatus = StatusError + msh.InstallErr = err + go msh.NotifyRemoteUpdate() +} + +func (msh *MShellProc) GetRemoteCopy() sstore.RemoteType { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return *msh.Remote +} + +func (msh *MShellProc) GetUName() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.UName +} + +func (msh *MShellProc) GetNumRunningCommands() int { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return len(msh.RunningCmds) +} + +func (msh *MShellProc) UpdateRemote(ctx context.Context, editMap map[string]interface{}) error { + msh.Lock.Lock() + defer msh.Lock.Unlock() + updatedRemote, err := sstore.UpdateRemote(ctx, msh.Remote.RemoteId, editMap) + if err != nil { + return err + } + if updatedRemote == nil { + return fmt.Errorf("no remote returned from UpdateRemote") + } + msh.Remote = updatedRemote + go msh.NotifyRemoteUpdate() + return nil +} + +func (msh *MShellProc) Disconnect(force bool) { + status := msh.GetStatus() + if status != StatusConnected && status != StatusConnecting { + msh.WriteToPtyBuffer("remote already disconnected (no action taken)\n") + return + } + numCommands := msh.GetNumRunningCommands() + if numCommands > 0 && !force { + msh.WriteToPtyBuffer("remote not disconnected, has %d running commands. use force=1 to force disconnection\n", numCommands) + return + } + msh.Lock.Lock() + defer msh.Lock.Unlock() + if msh.ServerProc != nil { + msh.ServerProc.Close() + } + if msh.MakeClientCancelFn != nil { + msh.MakeClientCancelFn() + msh.MakeClientCancelFn = nil + } +} + +func (msh *MShellProc) CancelInstall() { + msh.Lock.Lock() + defer msh.Lock.Unlock() + if msh.InstallCancelFn != nil { + msh.InstallCancelFn() + msh.InstallCancelFn = nil + } +} + +func (msh *MShellProc) GetRemoteName() string { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Remote.GetName() +} + +func (msh *MShellProc) WriteToPtyBuffer(strFmt string, args ...interface{}) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + msh.writeToPtyBuffer_nolock(strFmt, args...) +} + +func (msh *MShellProc) writeToPtyBuffer_nolock(strFmt string, args ...interface{}) { + // inefficient string manipulation here and read of PtyBuffer, but these messages are rare, nbd + realStr := fmt.Sprintf(strFmt, args...) + if !strings.HasPrefix(realStr, "~") { + realStr = strings.ReplaceAll(realStr, "\n", "\r\n") + if !strings.HasSuffix(realStr, "\r\n") { + realStr = realStr + "\r\n" + } + if strings.HasPrefix(realStr, "*") { + realStr = "\033[0m\033[31mprompt>\033[0m " + realStr[1:] + } else { + realStr = "\033[0m\033[32mprompt>\033[0m " + realStr + } + barr := msh.PtyBuffer.Bytes() + if len(barr) > 0 && barr[len(barr)-1] != '\n' { + realStr = "\r\n" + realStr + } + } else { + realStr = realStr[1:] + } + curOffset := msh.PtyBuffer.TotalWritten() + data := []byte(realStr) + msh.PtyBuffer.Write(data) + sendRemotePtyUpdate(msh.Remote.RemoteId, curOffset, data) +} + +func sendRemotePtyUpdate(remoteId string, dataOffset int64, data []byte) { + data64 := base64.StdEncoding.EncodeToString(data) + update := &sstore.PtyDataUpdate{ + RemoteId: remoteId, + PtyPos: dataOffset, + PtyData64: data64, + PtyDataLen: int64(len(data)), + } + sstore.MainBus.SendUpdate(update) +} + +func (msh *MShellProc) isWaitingForPassword_nolock() bool { + barr := msh.PtyBuffer.Bytes() + if len(barr) == 0 { + return false + } + nlIdx := bytes.LastIndex(barr, []byte{'\n'}) + var lastLine string + if nlIdx == -1 { + lastLine = string(barr) + } else { + lastLine = string(barr[nlIdx+1:]) + } + pwIdx := strings.Index(lastLine, "assword") + return pwIdx != -1 +} + +func (msh *MShellProc) RunPtyReadLoop(cmdPty *os.File) { + buf := make([]byte, PtyReadBufSize) + var isWaiting bool + for { + n, readErr := cmdPty.Read(buf) + if readErr == io.EOF { + break + } + if readErr != nil { + msh.WriteToPtyBuffer("*error reading from controlling-pty: %v\n", readErr) + break + } + var newIsWaiting bool + msh.WithLock(func() { + curOffset := msh.PtyBuffer.TotalWritten() + msh.PtyBuffer.Write(buf[0:n]) + sendRemotePtyUpdate(msh.Remote.RemoteId, curOffset, buf[0:n]) + newIsWaiting = msh.isWaitingForPassword_nolock() + }) + if newIsWaiting != isWaiting { + isWaiting = newIsWaiting + go msh.NotifyRemoteUpdate() + } + } +} + +func (msh *MShellProc) WaitAndSendPassword(pw string) { + var numWaits int + for { + var isWaiting bool + var isConnecting bool + msh.WithLock(func() { + isWaiting = msh.isWaitingForPassword_nolock() + isConnecting = msh.Status == StatusConnecting + }) + if !isConnecting { + break + } + if !isWaiting { + numWaits = 0 + time.Sleep(100 * time.Millisecond) + continue + } + numWaits++ + if numWaits < 10 { + time.Sleep(100 * time.Millisecond) + } else { + // send password + msh.WithLock(func() { + if msh.ControllingPty == nil { + return + } + pwBytes := []byte(pw + "\r") + msh.writeToPtyBuffer_nolock("~[sent password]\r\n") + _, err := msh.ControllingPty.Write(pwBytes) + if err != nil { + msh.writeToPtyBuffer_nolock("*cannot write password to controlling pty: %v\n", err) + } + }) + break + } + } +} + +func (msh *MShellProc) RunInstall() { + remoteCopy := msh.GetRemoteCopy() + if remoteCopy.Archived { + msh.WriteToPtyBuffer("*error: cannot install on archived remote\n") + return + } + baseStatus := msh.GetStatus() + if baseStatus == StatusConnecting || baseStatus == StatusConnected { + msh.WriteToPtyBuffer("*error: cannot install on remote that is connected/connecting, disconnect to install\n") + return + } + curStatus := msh.GetInstallStatus() + if curStatus == StatusConnecting { + msh.WriteToPtyBuffer("*error: cannot install on remote that is already trying to install, cancel current install to try again\n") + return + } + msh.WriteToPtyBuffer("installing mshell %s to %s...\n", scbase.MShellVersion, remoteCopy.RemoteCanonicalName) + sshOpts := convertSSHOpts(remoteCopy.SSHOpts) + sshOpts.SSHErrorsToTty = true + cmdStr := shexec.MakeInstallCommandStr() + ecmd := sshOpts.MakeSSHExecCmd(cmdStr) + cmdPty, err := msh.addControllingTty(ecmd) + if err != nil { + statusErr := fmt.Errorf("cannot attach controlling tty to mshell install command: %w", err) + msh.setInstallErrorStatus(statusErr) + return + } + defer func() { + if len(ecmd.ExtraFiles) > 0 { + ecmd.ExtraFiles[len(ecmd.ExtraFiles)-1].Close() + } + cmdPty.Close() + }() + go msh.RunPtyReadLoop(cmdPty) + clientCtx, clientCancelFn := context.WithCancel(context.Background()) + defer clientCancelFn() + msh.WithLock(func() { + msh.InstallErr = nil + msh.InstallStatus = StatusConnecting + msh.InstallCancelFn = clientCancelFn + go msh.NotifyRemoteUpdate() + }) + msgFn := func(msg string) { + msh.WriteToPtyBuffer("%s", msg) + } + err = shexec.RunInstallFromCmd(clientCtx, ecmd, true, nil, scbase.MShellBinaryReader, msgFn) + if err == context.Canceled { + msh.WriteToPtyBuffer("*install canceled\n") + msh.WithLock(func() { + msh.InstallStatus = StatusDisconnected + go msh.NotifyRemoteUpdate() + }) + return + } + if err != nil { + statusErr := fmt.Errorf("install failed: %w", err) + msh.setInstallErrorStatus(statusErr) + return + } + var connectMode string + msh.WithLock(func() { + msh.InstallStatus = StatusDisconnected + msh.InstallCancelFn = nil + msh.NeedsMShellUpgrade = false + msh.Status = StatusDisconnected + msh.Err = nil + connectMode = msh.Remote.ConnectMode + }) + msh.WriteToPtyBuffer("successfully installed mshell %s to ~/.mshell\n", scbase.MShellVersion) + go msh.NotifyRemoteUpdate() + if connectMode == sstore.ConnectModeStartup || connectMode == sstore.ConnectModeAuto { + // the install was successful, and we don't have a manual connect mode, try to connect + go msh.Launch(true) + } + return +} + +func (msh *MShellProc) updateRemoteStateVars(ctx context.Context, remoteId string, initPk *packet.InitPacketType) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + stateVars := getStateVarsFromInitPk(initPk) + if stateVars == nil { + return + } + msh.Remote.StateVars = stateVars + err := sstore.UpdateRemoteStateVars(ctx, remoteId, stateVars) + if err != nil { + // ignore error, nothing to do + log.Printf("error updating remote statevars: %v\n", err) + } +} + +func getStateVarsFromInitPk(initPk *packet.InitPacketType) map[string]string { + if initPk == nil || initPk.NotFound { + return nil + } + rtn := make(map[string]string) + rtn["home"] = initPk.HomeDir + rtn["remoteuser"] = initPk.User + rtn["remotehost"] = initPk.HostName + rtn["remoteuname"] = initPk.UName + return rtn +} + +func (msh *MShellProc) ReInit(ctx context.Context) (*packet.InitPacketType, error) { + reinitPk := packet.MakeReInitPacket() + reinitPk.ReqId = uuid.New().String() + resp, err := msh.PacketRpcRaw(ctx, reinitPk) + if err != nil { + return nil, err + } + if resp == nil { + return nil, fmt.Errorf("no response") + } + initPk, ok := resp.(*packet.InitPacketType) + if !ok { + return nil, fmt.Errorf("invalid reinit response (not an initpacket): %T", resp) + } + if initPk.State == nil { + return nil, fmt.Errorf("invalid reinit response initpk does not contain remote state") + } + hval := initPk.State.GetHashVal(false) + sstore.StoreStateBase(ctx, initPk.State) + msh.WithLock(func() { + msh.CurrentState = hval + msh.StateMap[hval] = initPk.State + }) + msh.updateRemoteStateVars(ctx, msh.RemoteId, initPk) + return initPk, nil +} + +func (msh *MShellProc) StreamFile(ctx context.Context, streamPk *packet.StreamFilePacketType) (*packet.RpcResponseIter, error) { + return msh.PacketRpcIter(ctx, streamPk) +} + +func addScVarsToState(state *packet.ShellState) *packet.ShellState { + if state == nil { + return nil + } + rtn := *state + envMap := shexec.DeclMapFromState(&rtn) + envMap["PROMPT"] = &shexec.DeclareDeclType{Name: "PROMPT", Value: "1", Args: "x"} + envMap["PROMPT_VERSION"] = &shexec.DeclareDeclType{Name: "PROMPT_VERSION", Value: scbase.PromptVersion, Args: "x"} + rtn.ShellVars = shexec.SerializeDeclMap(envMap) + return &rtn +} + +func stripScVarsFromState(state *packet.ShellState) *packet.ShellState { + if state == nil { + return nil + } + rtn := *state + rtn.HashVal = "" + envMap := shexec.DeclMapFromState(&rtn) + delete(envMap, "PROMPT") + delete(envMap, "PROMPT_VERSION") + rtn.ShellVars = shexec.SerializeDeclMap(envMap) + return &rtn +} + +func stripScVarsFromStateDiff(stateDiff *packet.ShellStateDiff) *packet.ShellStateDiff { + if stateDiff == nil || len(stateDiff.VarsDiff) == 0 { + return stateDiff + } + rtn := *stateDiff + rtn.HashVal = "" + var mapDiff statediff.MapDiffType + err := mapDiff.Decode(stateDiff.VarsDiff) + if err != nil { + log.Printf("error decoding statediff in stripScVarsFromStateDiff: %v\n", err) + return stateDiff + } + delete(mapDiff.ToAdd, "PROMPT") + delete(mapDiff.ToAdd, "PROMPT_VERSION") + rtn.VarsDiff = mapDiff.Encode() + return &rtn +} + +func (msh *MShellProc) Launch(interactive bool) { + remoteCopy := msh.GetRemoteCopy() + if remoteCopy.Archived { + msh.WriteToPtyBuffer("cannot launch archived remote\n") + return + } + curStatus := msh.GetStatus() + if curStatus == StatusConnected { + msh.WriteToPtyBuffer("remote is already connected (no action taken)\n") + return + } + if curStatus == StatusConnecting { + msh.WriteToPtyBuffer("remote is already connecting, disconnect before trying to connect again\n") + return + } + istatus := msh.GetInstallStatus() + if istatus == StatusConnecting { + msh.WriteToPtyBuffer("remote is trying to install, cancel install before trying to connect again\n") + return + } + if remoteCopy.SSHOpts.SSHPort != 0 && remoteCopy.SSHOpts.SSHPort != 22 { + msh.WriteToPtyBuffer("connecting to %s (port %d)...\n", remoteCopy.RemoteCanonicalName, remoteCopy.SSHOpts.SSHPort) + } else { + msh.WriteToPtyBuffer("connecting to %s...\n", remoteCopy.RemoteCanonicalName) + } + sshOpts := convertSSHOpts(remoteCopy.SSHOpts) + sshOpts.SSHErrorsToTty = true + if remoteCopy.ConnectMode != sstore.ConnectModeManual && remoteCopy.SSHOpts.SSHPassword == "" && !interactive { + sshOpts.BatchMode = true + } + var cmdStr string + if sshOpts.SSHHost == "" && remoteCopy.Local { + var err error + cmdStr, err = MakeLocalMShellCommandStr(remoteCopy.IsSudo()) + if err != nil { + msh.WriteToPtyBuffer("*error, cannot find local mshell binary: %v\n", err) + return + } + log.Printf("local mshell binary: %s\n", cmdStr) + } else { + cmdStr = MakeServerCommandStr() + } + ecmd := sshOpts.MakeSSHExecCmd(cmdStr) + cmdPty, err := msh.addControllingTty(ecmd) + if err != nil { + statusErr := fmt.Errorf("cannot attach controlling tty to mshell command: %w", err) + msh.WriteToPtyBuffer("*error, %s\n", statusErr.Error()) + msh.setErrorStatus(statusErr) + return + } + defer func() { + if len(ecmd.ExtraFiles) > 0 { + ecmd.ExtraFiles[len(ecmd.ExtraFiles)-1].Close() + } + }() + go msh.RunPtyReadLoop(cmdPty) + if remoteCopy.SSHOpts.SSHPassword != "" { + go msh.WaitAndSendPassword(remoteCopy.SSHOpts.SSHPassword) + } + makeClientCtx, makeClientCancelFn := context.WithCancel(context.Background()) + defer makeClientCancelFn() + msh.WithLock(func() { + msh.Err = nil + msh.ErrNoInitPk = false + msh.Status = StatusConnecting + msh.MakeClientCancelFn = makeClientCancelFn + deadlineTime := time.Now().Add(RemoteConnectTimeout) + msh.MakeClientDeadline = &deadlineTime + go msh.NotifyRemoteUpdate() + }) + go msh.watchClientDeadlineTime() + cproc, initPk, err := shexec.MakeClientProc(makeClientCtx, ecmd) + // TODO check if initPk.State is not nil + var mshellVersion string + var stateBaseHash string + var hitDeadline bool + msh.WithLock(func() { + msh.MakeClientCancelFn = nil + if time.Now().After(*msh.MakeClientDeadline) { + hitDeadline = true + } + msh.MakeClientDeadline = nil + if initPk == nil { + msh.ErrNoInitPk = true + } + if initPk != nil { + msh.UName = initPk.UName + mshellVersion = initPk.Version + if semver.Compare(mshellVersion, scbase.MShellVersion) < 0 { + // only set NeedsMShellUpgrade if we got an InitPk + msh.NeedsMShellUpgrade = true + } + } + if initPk != nil && initPk.State != nil { + hval := initPk.State.GetHashVal(false) + msh.CurrentState = hval + msh.StateMap[hval] = initPk.State + sstore.StoreStateBase(context.Background(), initPk.State) + stateBaseHash = hval + } else { + msh.CurrentState = "" + } + // no notify here, because we'll call notify in either case below + }) + if err == context.Canceled { + if hitDeadline { + msh.WriteToPtyBuffer("*connect timeout\n") + msh.setErrorStatus(errors.New("connect timeout")) + } else { + msh.WriteToPtyBuffer("*forced disconnection\n") + msh.WithLock(func() { + msh.Status = StatusDisconnected + go msh.NotifyRemoteUpdate() + }) + } + return + } + if err == nil && semver.MajorMinor(mshellVersion) != semver.MajorMinor(scbase.MShellVersion) { + err = fmt.Errorf("mshell version is not compatible current=%s remote=%s", scbase.MShellVersion, mshellVersion) + } + if err != nil { + msh.setErrorStatus(err) + msh.WriteToPtyBuffer("*error connecting to remote: %v\n", err) + go msh.tryAutoInstall() + return + } + msh.updateRemoteStateVars(context.Background(), msh.RemoteId, initPk) + msh.WriteToPtyBuffer("connected state:%s\n", stateBaseHash) + msh.WithLock(func() { + msh.ServerProc = cproc + msh.Status = StatusConnected + go msh.NotifyRemoteUpdate() + }) + go func() { + exitErr := cproc.Cmd.Wait() + exitCode := shexec.GetExitCode(exitErr) + msh.WithLock(func() { + if msh.Status == StatusConnected || msh.Status == StatusConnecting { + msh.Status = StatusDisconnected + go msh.NotifyRemoteUpdate() + } + }) + msh.WriteToPtyBuffer("*disconnected exitcode=%d\n", exitCode) + }() + go msh.ProcessPackets() + return +} + +func (msh *MShellProc) IsConnected() bool { + msh.Lock.Lock() + defer msh.Lock.Unlock() + return msh.Status == StatusConnected +} + +func replaceHomePath(pathStr string, homeDir string) string { + if homeDir == "" { + return pathStr + } + if pathStr == homeDir { + return "~" + } + if strings.HasPrefix(pathStr, homeDir+"/") { + return "~" + pathStr[len(homeDir):] + } + return pathStr +} + +func (state RemoteRuntimeState) ExpandHomeDir(pathStr string) (string, error) { + if pathStr != "~" && !strings.HasPrefix(pathStr, "~/") { + return pathStr, nil + } + homeDir := state.RemoteVars["home"] + if homeDir == "" { + return "", fmt.Errorf("remote does not have HOME set, cannot do ~ expansion") + } + if pathStr == "~" { + return homeDir, nil + } + return path.Join(homeDir, pathStr[2:]), nil +} + +func (msh *MShellProc) IsCmdRunning(ck base.CommandKey) bool { + msh.Lock.Lock() + defer msh.Lock.Unlock() + for runningCk, _ := range msh.RunningCmds { + if runningCk == ck { + return true + } + } + return false +} + +func (msh *MShellProc) SendInput(dataPk *packet.DataPacketType) error { + if !msh.IsConnected() { + return fmt.Errorf("remote is not connected, cannot send input") + } + if !msh.IsCmdRunning(dataPk.CK) { + return fmt.Errorf("cannot send input, cmd is not running") + } + return msh.ServerProc.Input.SendPacket(dataPk) +} + +func (msh *MShellProc) SendSpecialInput(siPk *packet.SpecialInputPacketType) error { + if !msh.IsConnected() { + return fmt.Errorf("remote is not connected, cannot send input") + } + if !msh.IsCmdRunning(siPk.CK) { + return fmt.Errorf("cannot send input, cmd is not running") + } + return msh.ServerProc.Input.SendPacket(siPk) +} + +func (msh *MShellProc) SendFileData(dataPk *packet.FileDataPacketType) error { + if !msh.IsConnected() { + return fmt.Errorf("remote is not connected, cannot send input") + } + return msh.ServerProc.Input.SendPacket(dataPk) +} + +func makeTermOpts(runPk *packet.RunPacketType) sstore.TermOpts { + return sstore.TermOpts{Rows: int64(runPk.TermOpts.Rows), Cols: int64(runPk.TermOpts.Cols), FlexRows: true, MaxPtySize: DefaultMaxPtySize} +} + +// returns (ok, currentPSC) +func (msh *MShellProc) testAndSetPendingStateCmd(riName string, newCK *base.CommandKey) (bool, *base.CommandKey) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + ck, found := msh.PendingStateCmds[riName] + if found { + return false, &ck + } + if newCK != nil { + msh.PendingStateCmds[riName] = *newCK + } + return true, nil +} + +func (msh *MShellProc) removePendingStateCmd(riName string, ck base.CommandKey) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + existingCK, found := msh.PendingStateCmds[riName] + if !found { + return + } + if existingCK == ck { + delete(msh.PendingStateCmds, riName) + } +} + +// returns (cmdtype, allow-updates-callback, err) +func RunCommand(ctx context.Context, sessionId string, screenId string, remotePtr sstore.RemotePtrType, runPacket *packet.RunPacketType) (rtnCmd *sstore.CmdType, rtnCallback func(), rtnErr error) { + rct := RunCmdType{ + SessionId: sessionId, + ScreenId: screenId, + RemotePtr: remotePtr, + RunPacket: runPacket, + } + if remotePtr.OwnerId != "" { + return nil, nil, fmt.Errorf("cannot run command against another user's remote '%s'", remotePtr.MakeFullRemoteRef()) + } + if screenId != runPacket.CK.GetGroupId() { + return nil, nil, fmt.Errorf("run commands screenids do not match") + } + msh := GetRemoteById(remotePtr.RemoteId) + if msh == nil { + return nil, nil, fmt.Errorf("no remote id=%s found", remotePtr.RemoteId) + } + if !msh.IsConnected() { + return nil, nil, fmt.Errorf("remote '%s' is not connected", remotePtr.RemoteId) + } + if runPacket.State != nil { + return nil, nil, fmt.Errorf("runPacket.State should not be set, it is set in RunCommand") + } + var newPSC *base.CommandKey + if runPacket.ReturnState { + newPSC = &runPacket.CK + } + ok, existingPSC := msh.testAndSetPendingStateCmd(remotePtr.Name, newPSC) + if !ok { + line, _, err := sstore.GetLineCmdByLineId(ctx, screenId, existingPSC.GetCmdId()) + if err != nil { + return nil, nil, fmt.Errorf("cannot run command while a stateful command is still running: %v", err) + } + if line == nil { + return nil, nil, fmt.Errorf("cannot run command while a stateful command is still running %s", *existingPSC) + } + return nil, nil, fmt.Errorf("cannot run command while a stateful command (linenum=%d) is still running", line.LineNum) + } + startCmdWait(runPacket.CK) + defer func() { + if rtnErr != nil { + removeCmdWait(runPacket.CK) + if newPSC != nil { + msh.removePendingStateCmd(remotePtr.Name, *newPSC) + } + } + }() + // get current remote-instance state + statePtr, err := sstore.GetRemoteStatePtr(ctx, sessionId, screenId, remotePtr) + if err != nil { + return nil, nil, fmt.Errorf("cannot get current remote stateptr: %w", err) + } + if statePtr == nil { + statePtr = msh.GetDefaultStatePtr() + } + if statePtr == nil { + return nil, nil, fmt.Errorf("cannot run command, no valid remote stateptr") + } + currentState, err := sstore.GetFullState(ctx, *statePtr) + if err != nil || currentState == nil { + return nil, nil, fmt.Errorf("cannot get current remote state: %w", err) + } + runPacket.State = addScVarsToState(currentState) + runPacket.StateComplete = true + msh.ServerProc.Output.RegisterRpc(runPacket.ReqId) + err = shexec.SendRunPacketAndRunData(ctx, msh.ServerProc.Input, runPacket) + if err != nil { + return nil, nil, fmt.Errorf("sending run packet to remote: %w", err) + } + rtnPk := msh.ServerProc.Output.WaitForResponse(ctx, runPacket.ReqId) + if rtnPk == nil { + return nil, nil, ctx.Err() + } + startPk, ok := rtnPk.(*packet.CmdStartPacketType) + if !ok { + respPk, ok := rtnPk.(*packet.ResponsePacketType) + if !ok { + return nil, nil, fmt.Errorf("invalid response received from server for run packet: %s", packet.AsString(rtnPk)) + } + if respPk.Error != "" { + return nil, nil, errors.New(respPk.Error) + } + return nil, nil, fmt.Errorf("invalid response received from server for run packet: %s", packet.AsString(rtnPk)) + } + status := sstore.CmdStatusRunning + if runPacket.Detached { + status = sstore.CmdStatusDetached + } + cmd := &sstore.CmdType{ + ScreenId: runPacket.CK.GetGroupId(), + LineId: runPacket.CK.GetCmdId(), + CmdStr: runPacket.Command, + RawCmdStr: runPacket.Command, + Remote: remotePtr, + FeState: sstore.FeStateFromShellState(currentState), + StatePtr: *statePtr, + TermOpts: makeTermOpts(runPacket), + Status: status, + CmdPid: startPk.Pid, + RemotePid: startPk.MShellPid, + ExitCode: 0, + DurationMs: 0, + RunOut: nil, + RtnState: runPacket.ReturnState, + } + err = sstore.CreateCmdPtyFile(ctx, cmd.ScreenId, cmd.LineId, cmd.TermOpts.MaxPtySize) + if err != nil { + // TODO the cmd is running, so this is a tricky error to handle + return nil, nil, fmt.Errorf("cannot create local ptyout file for running command: %v", err) + } + msh.AddRunningCmd(rct) + return cmd, func() { removeCmdWait(runPacket.CK) }, nil +} + +func (msh *MShellProc) AddWaitingCmd(rct RunCmdType) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + msh.WaitingCmds = append(msh.WaitingCmds, rct) +} + +func (msh *MShellProc) reExecSingle(rct RunCmdType) { + // TODO fixme + ctx, cancelFn := context.WithTimeout(context.Background(), 15*time.Second) + defer cancelFn() + _, callback, _ := RunCommand(ctx, rct.SessionId, rct.ScreenId, rct.RemotePtr, rct.RunPacket) + if callback != nil { + defer callback() + } +} + +func (msh *MShellProc) ReExecWaitingCmds() { + msh.Lock.Lock() + defer msh.Lock.Unlock() + for len(msh.WaitingCmds) > 0 { + rct := msh.WaitingCmds[0] + go msh.reExecSingle(rct) + if rct.RunPacket.ReturnState { + break + } + } + if len(msh.WaitingCmds) == 0 { + msh.WaitingCmds = nil + } +} + +func (msh *MShellProc) AddRunningCmd(rct RunCmdType) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + msh.RunningCmds[rct.RunPacket.CK] = rct +} + +func (msh *MShellProc) GetRunningCmd(ck base.CommandKey) *RunCmdType { + msh.Lock.Lock() + defer msh.Lock.Unlock() + rct, found := msh.RunningCmds[ck] + if !found { + return nil + } + return &rct +} + +func (msh *MShellProc) RemoveRunningCmd(ck base.CommandKey) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + delete(msh.RunningCmds, ck) + for name, pendingCk := range msh.PendingStateCmds { + if pendingCk == ck { + delete(msh.PendingStateCmds, name) + } + } +} + +func (msh *MShellProc) PacketRpcIter(ctx context.Context, pk packet.RpcPacketType) (*packet.RpcResponseIter, error) { + if !msh.IsConnected() { + return nil, fmt.Errorf("remote is not connected") + } + if pk == nil { + return nil, fmt.Errorf("PacketRpc passed nil packet") + } + reqId := pk.GetReqId() + msh.ServerProc.Output.RegisterRpc(reqId) + err := msh.ServerProc.Input.SendPacketCtx(ctx, pk) + if err != nil { + return nil, err + } + return msh.ServerProc.Output.GetResponseIter(reqId), nil +} + +func (msh *MShellProc) PacketRpcRaw(ctx context.Context, pk packet.RpcPacketType) (packet.RpcResponsePacketType, error) { + if !msh.IsConnected() { + return nil, fmt.Errorf("remote is not connected") + } + if pk == nil { + return nil, fmt.Errorf("PacketRpc passed nil packet") + } + reqId := pk.GetReqId() + msh.ServerProc.Output.RegisterRpc(reqId) + defer msh.ServerProc.Output.UnRegisterRpc(reqId) + err := msh.ServerProc.Input.SendPacketCtx(ctx, pk) + if err != nil { + return nil, err + } + rtnPk := msh.ServerProc.Output.WaitForResponse(ctx, reqId) + if rtnPk == nil { + return nil, ctx.Err() + } + return rtnPk, nil +} + +func (msh *MShellProc) PacketRpc(ctx context.Context, pk packet.RpcPacketType) (*packet.ResponsePacketType, error) { + rtnPk, err := msh.PacketRpcRaw(ctx, pk) + if err != nil { + return nil, err + } + if respPk, ok := rtnPk.(*packet.ResponsePacketType); ok { + return respPk, nil + } + return nil, fmt.Errorf("invalid response packet received: %s", packet.AsString(rtnPk)) +} + +func (msh *MShellProc) WithLock(fn func()) { + msh.Lock.Lock() + defer msh.Lock.Unlock() + fn() +} + +func makeDataAckPacket(ck base.CommandKey, fdNum int, ackLen int, err error) *packet.DataAckPacketType { + ack := packet.MakeDataAckPacket() + ack.CK = ck + ack.FdNum = fdNum + ack.AckLen = ackLen + if err != nil { + ack.Error = err.Error() + } + return ack +} + +func (msh *MShellProc) notifyHangups_nolock() { + for ck, _ := range msh.RunningCmds { + cmd, err := sstore.GetCmdByScreenId(context.Background(), ck.GetGroupId(), ck.GetCmdId()) + if err != nil { + continue + } + update := &sstore.ModelUpdate{Cmd: cmd} + sstore.MainBus.SendScreenUpdate(ck.GetGroupId(), update) + } + msh.RunningCmds = make(map[base.CommandKey]RunCmdType) +} + +func (msh *MShellProc) handleCmdDonePacket(donePk *packet.CmdDonePacketType) { + // this will remove from RunningCmds and from PendingStateCmds + defer msh.RemoveRunningCmd(donePk.CK) + if donePk.FinalState != nil { + donePk.FinalState = stripScVarsFromState(donePk.FinalState) + } + if donePk.FinalStateDiff != nil { + donePk.FinalStateDiff = stripScVarsFromStateDiff(donePk.FinalStateDiff) + } + update, err := sstore.UpdateCmdDoneInfo(context.Background(), donePk.CK, donePk, sstore.CmdStatusDone) + if err != nil { + msh.WriteToPtyBuffer("*error updating cmddone: %v\n", err) + return + } + screen, err := sstore.UpdateScreenFocusForDoneCmd(context.Background(), donePk.CK.GetGroupId(), donePk.CK.GetCmdId()) + if err != nil { + msh.WriteToPtyBuffer("*error trying to update screen focus type: %v\n", err) + // fall-through (nothing to do) + } + if screen != nil { + update.Screens = []*sstore.ScreenType{screen} + } + rct := msh.GetRunningCmd(donePk.CK) + var statePtr *sstore.ShellStatePtr + if donePk.FinalState != nil && rct != nil { + feState := sstore.FeStateFromShellState(donePk.FinalState) + remoteInst, err := sstore.UpdateRemoteState(context.Background(), rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, donePk.FinalState, nil) + if err != nil { + msh.WriteToPtyBuffer("*error trying to update remotestate: %v\n", err) + // fall-through (nothing to do) + } + if remoteInst != nil { + update.Sessions = sstore.MakeSessionsUpdateForRemote(rct.SessionId, remoteInst) + } + statePtr = &sstore.ShellStatePtr{BaseHash: donePk.FinalState.GetHashVal(false)} + } else if donePk.FinalStateDiff != nil && rct != nil { + feState, err := msh.getFeStateFromDiff(donePk.FinalStateDiff) + if err != nil { + msh.WriteToPtyBuffer("*error trying to update remotestate: %v\n", err) + // fall-through (nothing to do) + } else { + remoteInst, err := sstore.UpdateRemoteState(context.Background(), rct.SessionId, rct.ScreenId, rct.RemotePtr, feState, nil, donePk.FinalStateDiff) + if err != nil { + msh.WriteToPtyBuffer("*error trying to update remotestate: %v\n", err) + // fall-through (nothing to do) + } + if remoteInst != nil { + update.Sessions = sstore.MakeSessionsUpdateForRemote(rct.SessionId, remoteInst) + } + diffHashArr := append(([]string)(nil), donePk.FinalStateDiff.DiffHashArr...) + diffHashArr = append(diffHashArr, donePk.FinalStateDiff.GetHashVal(false)) + statePtr = &sstore.ShellStatePtr{BaseHash: donePk.FinalStateDiff.BaseHash, DiffHashArr: diffHashArr} + } + } + if statePtr != nil { + err = sstore.UpdateCmdRtnState(context.Background(), donePk.CK, *statePtr) + if err != nil { + msh.WriteToPtyBuffer("*error trying to update cmd rtnstate: %v\n", err) + // fall-through (nothing to do) + } + } + sstore.MainBus.SendUpdate(update) + return +} + +func (msh *MShellProc) handleCmdFinalPacket(finalPk *packet.CmdFinalPacketType) { + defer msh.RemoveRunningCmd(finalPk.CK) + rtnCmd, err := sstore.GetCmdByScreenId(context.Background(), finalPk.CK.GetGroupId(), finalPk.CK.GetCmdId()) + if err != nil { + log.Printf("error calling GetCmdById in handleCmdFinalPacket: %v\n", err) + return + } + if rtnCmd == nil || rtnCmd.DoneTs > 0 { + return + } + log.Printf("finalpk %s (hangup): %s\n", finalPk.CK, finalPk.Error) + screen, err := sstore.HangupCmd(context.Background(), finalPk.CK) + if err != nil { + log.Printf("error in hangup-cmd in handleCmdFinalPacket: %v\n", err) + return + } + rtnCmd, err = sstore.GetCmdByScreenId(context.Background(), finalPk.CK.GetGroupId(), finalPk.CK.GetCmdId()) + if err != nil { + log.Printf("error getting cmd(2) in handleCmdFinalPacket: %v\n", err) + return + } + if rtnCmd == nil { + log.Printf("error getting cmd(2) in handleCmdFinalPacket (not found)\n") + return + } + update := &sstore.ModelUpdate{Cmd: rtnCmd} + if screen != nil { + update.Screens = []*sstore.ScreenType{screen} + } + sstore.MainBus.SendUpdate(update) +} + +// TODO notify FE about cmd errors +func (msh *MShellProc) handleCmdErrorPacket(errPk *packet.CmdErrorPacketType) { + err := sstore.AppendCmdErrorPk(context.Background(), errPk) + if err != nil { + msh.WriteToPtyBuffer("cmderr> [remote %s] [error] adding cmderr: %v\n", msh.GetRemoteName(), err) + return + } + return +} + +func (msh *MShellProc) handleDataPacket(dataPk *packet.DataPacketType, dataPosMap map[base.CommandKey]int64) { + realData, err := base64.StdEncoding.DecodeString(dataPk.Data64) + if err != nil { + ack := makeDataAckPacket(dataPk.CK, dataPk.FdNum, 0, err) + msh.ServerProc.Input.SendPacket(ack) + return + } + var ack *packet.DataAckPacketType + if len(realData) > 0 { + dataPos := dataPosMap[dataPk.CK] + rcmd := msh.GetRunningCmd(dataPk.CK) + update, err := sstore.AppendToCmdPtyBlob(context.Background(), rcmd.ScreenId, dataPk.CK.GetCmdId(), realData, dataPos) + if err != nil { + ack = makeDataAckPacket(dataPk.CK, dataPk.FdNum, 0, err) + } else { + ack = makeDataAckPacket(dataPk.CK, dataPk.FdNum, len(realData), nil) + } + dataPosMap[dataPk.CK] += int64(len(realData)) + if update != nil { + sstore.MainBus.SendScreenUpdate(dataPk.CK.GetGroupId(), update) + } + } + if ack != nil { + msh.ServerProc.Input.SendPacket(ack) + } + // log.Printf("data %s fd=%d len=%d eof=%v err=%v\n", dataPk.CK, dataPk.FdNum, len(realData), dataPk.Eof, dataPk.Error) +} + +func (msh *MShellProc) makeHandleDataPacketClosure(dataPk *packet.DataPacketType, dataPosMap map[base.CommandKey]int64) func() { + return func() { + msh.handleDataPacket(dataPk, dataPosMap) + } +} + +func (msh *MShellProc) makeHandleCmdDonePacketClosure(donePk *packet.CmdDonePacketType) func() { + return func() { + msh.handleCmdDonePacket(donePk) + } +} + +func (msh *MShellProc) makeHandleCmdFinalPacketClosure(finalPk *packet.CmdFinalPacketType) func() { + return func() { + msh.handleCmdFinalPacket(finalPk) + } +} + +func sendScreenUpdates(screens []*sstore.ScreenType) { + for _, screen := range screens { + sstore.MainBus.SendUpdate(&sstore.ModelUpdate{Screens: []*sstore.ScreenType{screen}}) + } +} + +func (msh *MShellProc) ProcessPackets() { + defer msh.WithLock(func() { + if msh.Status == StatusConnected { + msh.Status = StatusDisconnected + } + screens, err := sstore.HangupRunningCmdsByRemoteId(context.Background(), msh.Remote.RemoteId) + if err != nil { + msh.writeToPtyBuffer_nolock("error calling HUP on cmds %v\n", err) + } + msh.notifyHangups_nolock() + go msh.NotifyRemoteUpdate() + if len(screens) > 0 { + go sendScreenUpdates(screens) + } + }) + // TODO need to clean dataPosMap + dataPosMap := make(map[base.CommandKey]int64) + for pk := range msh.ServerProc.Output.MainCh { + if pk.GetType() == packet.DataPacketStr { + dataPk := pk.(*packet.DataPacketType) + runCmdUpdateFn(dataPk.CK, msh.makeHandleDataPacketClosure(dataPk, dataPosMap)) + continue + } + if pk.GetType() == packet.DataAckPacketStr { + // TODO process ack (need to keep track of buffer size for sending) + // this is low priority though since most input is coming from keyboard and won't overflow this buffer + continue + } + if pk.GetType() == packet.CmdDataPacketStr { + dataPacket := pk.(*packet.CmdDataPacketType) + msh.WriteToPtyBuffer("cmd-data> [remote %s] [%s] pty=%d run=%d\n", msh.GetRemoteName(), dataPacket.CK, dataPacket.PtyDataLen, dataPacket.RunDataLen) + continue + } + if pk.GetType() == packet.CmdDonePacketStr { + donePk := pk.(*packet.CmdDonePacketType) + runCmdUpdateFn(donePk.CK, msh.makeHandleCmdDonePacketClosure(donePk)) + continue + } + if pk.GetType() == packet.CmdFinalPacketStr { + finalPk := pk.(*packet.CmdFinalPacketType) + runCmdUpdateFn(finalPk.CK, msh.makeHandleCmdFinalPacketClosure(finalPk)) + continue + } + if pk.GetType() == packet.CmdErrorPacketStr { + msh.handleCmdErrorPacket(pk.(*packet.CmdErrorPacketType)) + continue + } + if pk.GetType() == packet.MessagePacketStr { + msgPacket := pk.(*packet.MessagePacketType) + msh.WriteToPtyBuffer("msg> [remote %s] [%s] %s\n", msh.GetRemoteName(), msgPacket.CK, msgPacket.Message) + continue + } + if pk.GetType() == packet.RawPacketStr { + rawPacket := pk.(*packet.RawPacketType) + msh.WriteToPtyBuffer("stderr> [remote %s] %s\n", msh.GetRemoteName(), rawPacket.Data) + continue + } + if pk.GetType() == packet.CmdStartPacketStr { + startPk := pk.(*packet.CmdStartPacketType) + msh.WriteToPtyBuffer("start> [remote %s] reqid=%s (%p)\n", msh.GetRemoteName(), startPk.RespId, msh.ServerProc.Output) + continue + } + msh.WriteToPtyBuffer("MSH> [remote %s] unhandled packet %s\n", msh.GetRemoteName(), packet.AsString(pk)) + } +} + +// returns number of chars (including braces) for brace-expr +func getBracedStr(runeStr []rune) int { + if len(runeStr) < 3 { + return 0 + } + if runeStr[0] != '{' { + return 0 + } + for i := 1; i < len(runeStr); i++ { + if runeStr[i] == '}' { + if i == 1 { // cannot have {} + return 0 + } + return i + 1 + } + } + return 0 +} + +func isDigit(r rune) bool { + return r >= '0' && r <= '9' // just check ascii digits (not unicode) +} + +func EvalPrompt(promptFmt string, vars map[string]string, state *packet.ShellState) string { + var buf bytes.Buffer + promptRunes := []rune(promptFmt) + for i := 0; i < len(promptRunes); i++ { + ch := promptRunes[i] + if ch == '\\' && i != len(promptRunes)-1 { + nextCh := promptRunes[i+1] + if nextCh == 'x' || nextCh == 'y' { + nr := getBracedStr(promptRunes[i+2:]) + if nr > 0 { + escCode := string(promptRunes[i+1 : i+1+nr+1]) // start at "x" or "y", extend nr+1 runes + escStr := evalPromptEsc(escCode, vars, state) + buf.WriteString(escStr) + i += nr + 1 + continue + } else { + buf.WriteRune(ch) // invalid escape, so just write ch and move on + continue + } + } else if isDigit(nextCh) { + if len(promptRunes) >= i+4 && isDigit(promptRunes[i+2]) && isDigit(promptRunes[i+3]) { + i += 3 + escStr := evalPromptEsc(string(promptRunes[i+1:i+4]), vars, state) + buf.WriteString(escStr) + continue + } else { + buf.WriteRune(ch) // invalid escape, so just write ch and move on + continue + } + } else { + i += 1 + escStr := evalPromptEsc(string(nextCh), vars, state) + buf.WriteString(escStr) + continue + } + } + buf.WriteRune(ch) + } + return buf.String() +} + +func evalPromptEsc(escCode string, vars map[string]string, state *packet.ShellState) string { + if strings.HasPrefix(escCode, "x{") && strings.HasSuffix(escCode, "}") { + varName := escCode[2 : len(escCode)-1] + return vars[varName] + } + if strings.HasPrefix(escCode, "y{") && strings.HasSuffix(escCode, "}") { + if state == nil { + return "" + } + varName := escCode[2 : len(escCode)-1] + varMap := shexec.ShellVarMapFromState(state) + return varMap[varName] + } + if escCode == "h" { + return vars["remoteshorthost"] + } + if escCode == "H" { + return vars["remotehost"] + } + if escCode == "s" { + return "mshell" + } + if escCode == "u" { + return vars["remoteuser"] + } + if escCode == "w" { + if state == nil { + return "?" + } + return replaceHomePath(state.Cwd, vars["home"]) + } + if escCode == "W" { + if state == nil { + return "?" + } + return path.Base(replaceHomePath(state.Cwd, vars["home"])) + } + if escCode == "$" { + if vars["remoteuser"] == "root" || vars["sudo"] == "1" { + return "#" + } else { + return "$" + } + } + if len(escCode) == 3 { + // \nnn escape + ival, err := strconv.ParseInt(escCode, 8, 32) + if err != nil { + return escCode + } + return string([]byte{byte(ival)}) + } + if escCode == "e" { + return "\033" + } + if escCode == "n" { + return "\n" + } + if escCode == "r" { + return "\r" + } + if escCode == "a" { + return "\007" + } + if escCode == "\\" { + return "\\" + } + if escCode == "[" { + return "" + } + if escCode == "]" { + return "" + } + + // we don't support date/time escapes (d, t, T, @), version escapes (v, V), cmd number (#, !), terminal device (l), jobs (j) + return "(" + escCode + ")" +} + +func (msh *MShellProc) getFullState(stateDiff *packet.ShellStateDiff) (*packet.ShellState, error) { + baseState := msh.GetStateByHash(stateDiff.BaseHash) + if baseState != nil && len(stateDiff.DiffHashArr) == 0 { + newState, err := shexec.ApplyShellStateDiff(*baseState, *stateDiff) + if err != nil { + return nil, err + } + return &newState, nil + } else { + fullState, err := sstore.GetFullState(context.Background(), sstore.ShellStatePtr{stateDiff.BaseHash, stateDiff.DiffHashArr}) + if err != nil { + return nil, err + } + newState, err := shexec.ApplyShellStateDiff(*fullState, *stateDiff) + return &newState, nil + } +} + +// internal func, first tries the StateMap, otherwise will fallback on sstore.GetFullState +func (msh *MShellProc) getFeStateFromDiff(stateDiff *packet.ShellStateDiff) (map[string]string, error) { + baseState := msh.GetStateByHash(stateDiff.BaseHash) + if baseState != nil && len(stateDiff.DiffHashArr) == 0 { + newState, err := shexec.ApplyShellStateDiff(*baseState, *stateDiff) + if err != nil { + return nil, err + } + return sstore.FeStateFromShellState(&newState), nil + } else { + fullState, err := sstore.GetFullState(context.Background(), sstore.ShellStatePtr{stateDiff.BaseHash, stateDiff.DiffHashArr}) + if err != nil { + return nil, err + } + newState, err := shexec.ApplyShellStateDiff(*fullState, *stateDiff) + if err != nil { + return nil, err + } + return sstore.FeStateFromShellState(&newState), nil + } +} + +func (msh *MShellProc) TryAutoConnect() error { + if msh.IsConnected() { + return nil + } + rcopy := msh.GetRemoteCopy() + if rcopy.ConnectMode == sstore.ConnectModeManual { + return nil + } + var err error + msh.WithLock(func() { + if msh.NumTryConnect > 5 { + err = fmt.Errorf("too many unsuccessful tries") + return + } + msh.NumTryConnect++ + }) + if err != nil { + return err + } + msh.Launch(false) + if !msh.IsConnected() { + return fmt.Errorf("error connecting") + } + return nil +} + +func (msh *MShellProc) GetDisplayName() string { + rcopy := msh.GetRemoteCopy() + return rcopy.GetName() +} diff --git a/wavesrv/pkg/remote/updatequeue.go b/wavesrv/pkg/remote/updatequeue.go new file mode 100644 index 00000000..c93ae6f6 --- /dev/null +++ b/wavesrv/pkg/remote/updatequeue.go @@ -0,0 +1,67 @@ +package remote + +import ( + "github.com/commandlinedev/apishell/pkg/base" +) + +func startCmdWait(ck base.CommandKey) { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + GlobalStore.CmdWaitMap[ck] = nil +} + +func pushCmdWaitIfRequired(ck base.CommandKey, fn func()) bool { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + fns, ok := GlobalStore.CmdWaitMap[ck] + if !ok { + return false + } + fns = append(fns, fn) + GlobalStore.CmdWaitMap[ck] = fns + return true +} + +func runCmdUpdateFn(ck base.CommandKey, fn func()) { + pushed := pushCmdWaitIfRequired(ck, fn) + if pushed { + return + } + fn() +} + +func runCmdWaitFns(ck base.CommandKey) { + for { + fn := removeFirstCmdWaitFn(ck) + if fn == nil { + break + } + fn() + } +} + +func removeFirstCmdWaitFn(ck base.CommandKey) func() { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + + fns := GlobalStore.CmdWaitMap[ck] + if len(fns) == 0 { + delete(GlobalStore.CmdWaitMap, ck) + return nil + } + fn := fns[0] + GlobalStore.CmdWaitMap[ck] = fns[1:] + return fn +} + +func removeCmdWait(ck base.CommandKey) { + GlobalStore.Lock.Lock() + defer GlobalStore.Lock.Unlock() + + fns := GlobalStore.CmdWaitMap[ck] + if len(fns) == 0 { + delete(GlobalStore.CmdWaitMap, ck) + return + } + go runCmdWaitFns(ck) +} diff --git a/wavesrv/pkg/rtnstate/rtnstate.go b/wavesrv/pkg/rtnstate/rtnstate.go new file mode 100644 index 00000000..2a27f0c0 --- /dev/null +++ b/wavesrv/pkg/rtnstate/rtnstate.go @@ -0,0 +1,196 @@ +package rtnstate + +import ( + "bytes" + "context" + "fmt" + "strings" + + "github.com/alessio/shellescape" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/apishell/pkg/simpleexpand" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/commandlinedev/prompt-server/pkg/utilfn" + "mvdan.cc/sh/v3/syntax" +) + +func parseAliasStmt(stmt *syntax.Stmt, sourceStr string) (string, string, error) { + cmd := stmt.Cmd + callExpr, ok := cmd.(*syntax.CallExpr) + if !ok { + return "", "", fmt.Errorf("wrong cmd type for alias") + } + if len(callExpr.Args) != 2 { + return "", "", fmt.Errorf("wrong number of words in alias expr wordslen=%d", len(callExpr.Args)) + } + firstWord := callExpr.Args[0] + if firstWord.Lit() != "alias" { + return "", "", fmt.Errorf("invalid alias cmd word (not 'alias')") + } + secondWord := callExpr.Args[1] + var ectx simpleexpand.SimpleExpandContext // no homedir, do not want ~ expansion + val, _ := simpleexpand.SimpleExpandWord(ectx, secondWord, sourceStr) + eqIdx := strings.Index(val, "=") + if eqIdx == -1 { + return "", "", fmt.Errorf("no '=' in alias definition") + } + return val[0:eqIdx], val[eqIdx+1:], nil +} + +func ParseAliases(aliases string) (map[string]string, error) { + r := strings.NewReader(aliases) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(r, "aliases") + if err != nil { + return nil, err + } + rtn := make(map[string]string) + for _, stmt := range file.Stmts { + aliasName, aliasVal, err := parseAliasStmt(stmt, aliases) + if err != nil { + // fmt.Printf("stmt-err: %v\n", err) + continue + } + if aliasName != "" { + rtn[aliasName] = aliasVal + } + } + return rtn, nil +} + +func parseFuncStmt(stmt *syntax.Stmt, source string) (string, string, error) { + cmd := stmt.Cmd + funcDecl, ok := cmd.(*syntax.FuncDecl) + if !ok { + return "", "", fmt.Errorf("cmd not FuncDecl") + } + name := funcDecl.Name.Value + // fmt.Printf("func: [%s]\n", name) + funcBody := funcDecl.Body + // fmt.Printf(" %d:%d\n", funcBody.Cmd.Pos().Offset(), funcBody.Cmd.End().Offset()) + bodyStr := source[funcBody.Cmd.Pos().Offset():funcBody.Cmd.End().Offset()] + // fmt.Printf("<<<\n%s\n>>>\n", bodyStr) + // fmt.Printf("\n") + return name, bodyStr, nil +} + +func ParseFuncs(funcs string) (map[string]string, error) { + r := strings.NewReader(funcs) + parser := syntax.NewParser(syntax.Variant(syntax.LangBash)) + file, err := parser.Parse(r, "funcs") + if err != nil { + return nil, err + } + rtn := make(map[string]string) + for _, stmt := range file.Stmts { + funcName, funcVal, err := parseFuncStmt(stmt, funcs) + if err != nil { + // TODO where to put parse errors + continue + } + if strings.HasPrefix(funcName, "_mshell_") { + continue + } + if funcName != "" { + rtn[funcName] = funcVal + } + } + return rtn, nil +} + +const MaxDiffKeyLen = 40 +const MaxDiffValLen = 50 + +var IgnoreVars = map[string]bool{"PROMPT": true, "PROMPT_VERSION": true, "MSHELL": true} + +func displayStateUpdateDiff(buf *bytes.Buffer, oldState packet.ShellState, newState packet.ShellState) { + if newState.Cwd != oldState.Cwd { + buf.WriteString(fmt.Sprintf("cwd %s\n", newState.Cwd)) + } + if !bytes.Equal(newState.ShellVars, oldState.ShellVars) { + newEnvMap := shexec.DeclMapFromState(&newState) + oldEnvMap := shexec.DeclMapFromState(&oldState) + for key, newVal := range newEnvMap { + if IgnoreVars[key] { + continue + } + oldVal, found := oldEnvMap[key] + if !found || !shexec.DeclsEqual(false, oldVal, newVal) { + var exportStr string + if newVal.IsExport() { + exportStr = "export " + } + buf.WriteString(fmt.Sprintf("%s%s=%s\n", exportStr, utilfn.EllipsisStr(key, MaxDiffKeyLen), utilfn.EllipsisStr(newVal.Value, MaxDiffValLen))) + } + } + for key, _ := range oldEnvMap { + if IgnoreVars[key] { + continue + } + _, found := newEnvMap[key] + if !found { + buf.WriteString(fmt.Sprintf("unset %s\n", utilfn.EllipsisStr(key, MaxDiffKeyLen))) + } + } + } + if newState.Aliases != oldState.Aliases { + newAliasMap, _ := ParseAliases(newState.Aliases) + oldAliasMap, _ := ParseAliases(oldState.Aliases) + for aliasName, newAliasVal := range newAliasMap { + oldAliasVal, found := oldAliasMap[aliasName] + if !found || newAliasVal != oldAliasVal { + buf.WriteString(fmt.Sprintf("alias %s\n", utilfn.EllipsisStr(shellescape.Quote(aliasName), MaxDiffKeyLen))) + } + } + for aliasName, _ := range oldAliasMap { + _, found := newAliasMap[aliasName] + if !found { + buf.WriteString(fmt.Sprintf("unalias %s\n", utilfn.EllipsisStr(shellescape.Quote(aliasName), MaxDiffKeyLen))) + } + } + } + if newState.Funcs != oldState.Funcs { + newFuncMap, _ := ParseFuncs(newState.Funcs) + oldFuncMap, _ := ParseFuncs(oldState.Funcs) + for funcName, newFuncVal := range newFuncMap { + oldFuncVal, found := oldFuncMap[funcName] + if !found || newFuncVal != oldFuncVal { + buf.WriteString(fmt.Sprintf("function %s\n", utilfn.EllipsisStr(shellescape.Quote(funcName), MaxDiffKeyLen))) + } + } + for funcName, _ := range oldFuncMap { + _, found := newFuncMap[funcName] + if !found { + buf.WriteString(fmt.Sprintf("unset -f %s\n", utilfn.EllipsisStr(shellescape.Quote(funcName), MaxDiffKeyLen))) + } + } + } +} + +func GetRtnStateDiff(ctx context.Context, screenId string, lineId string) ([]byte, error) { + cmd, err := sstore.GetCmdByScreenId(ctx, screenId, lineId) + if err != nil { + return nil, err + } + if cmd == nil { + return nil, nil + } + if !cmd.RtnState { + return nil, nil + } + if cmd.RtnStatePtr.IsEmpty() { + return nil, nil + } + var outputBytes bytes.Buffer + initialState, err := sstore.GetFullState(ctx, cmd.StatePtr) + if err != nil { + return nil, fmt.Errorf("getting initial full state: %v", err) + } + rtnState, err := sstore.GetFullState(ctx, cmd.RtnStatePtr) + if err != nil { + return nil, fmt.Errorf("getting rtn full state: %v", err) + } + displayStateUpdateDiff(&outputBytes, *initialState, *rtnState) + return outputBytes.Bytes(), nil +} diff --git a/wavesrv/pkg/scbase/scbase.go b/wavesrv/pkg/scbase/scbase.go new file mode 100644 index 00000000..7907746e --- /dev/null +++ b/wavesrv/pkg/scbase/scbase.go @@ -0,0 +1,407 @@ +package scbase + +import ( + "context" + "errors" + "fmt" + "io" + "io/fs" + "log" + "os" + "os/exec" + "os/user" + "path" + "regexp" + "runtime" + "strconv" + "strings" + "sync" + "time" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/google/uuid" + "golang.org/x/mod/semver" + "golang.org/x/sys/unix" +) + +const HomeVarName = "HOME" +const PromptHomeVarName = "PROMPT_HOME" +const PromptDevVarName = "PROMPT_DEV" +const SessionsDirBaseName = "sessions" +const ScreensDirBaseName = "screens" +const PromptLockFile = "prompt.lock" +const PromptDirName = "prompt" +const PromptDevDirName = "prompt-dev" +const PromptAppPathVarName = "PROMPT_APP_PATH" +const PromptVersion = "v0.4.0" +const PromptAuthKeyFileName = "prompt.authkey" +const MShellVersion = "v0.3.0" +const DefaultMacOSShell = "/bin/bash" + +var SessionDirCache = make(map[string]string) +var ScreenDirCache = make(map[string]string) +var BaseLock = &sync.Mutex{} +var BuildTime = "-" + +func IsDevMode() bool { + pdev := os.Getenv(PromptDevVarName) + return pdev != "" +} + +// must match js +func GetPromptHomeDir() string { + scHome := os.Getenv(PromptHomeVarName) + if scHome == "" { + homeVar := os.Getenv(HomeVarName) + if homeVar == "" { + homeVar = "/" + } + pdev := os.Getenv(PromptDevVarName) + if pdev != "" { + scHome = path.Join(homeVar, PromptDevDirName) + } else { + scHome = path.Join(homeVar, PromptDirName) + } + + } + return scHome +} + +func MShellBinaryDir() string { + appPath := os.Getenv(PromptAppPathVarName) + if appPath == "" { + appPath = "." + } + if IsDevMode() { + return path.Join(appPath, "dev-bin") + } + return path.Join(appPath, "bin", "mshell") +} + +func MShellBinaryPath(version string, goos string, goarch string) (string, error) { + if !base.ValidGoArch(goos, goarch) { + return "", fmt.Errorf("invalid goos/goarch combination: %s/%s", goos, goarch) + } + binaryDir := MShellBinaryDir() + versionStr := semver.MajorMinor(version) + if versionStr == "" { + return "", fmt.Errorf("invalid mshell version: %q", version) + } + fileName := fmt.Sprintf("mshell-%s-%s.%s", versionStr, goos, goarch) + fullFileName := path.Join(binaryDir, fileName) + return fullFileName, nil +} + +func LocalMShellBinaryPath() (string, error) { + return MShellBinaryPath(MShellVersion, runtime.GOOS, runtime.GOARCH) +} + +func MShellBinaryReader(version string, goos string, goarch string) (io.ReadCloser, error) { + mshellPath, err := MShellBinaryPath(version, goos, goarch) + if err != nil { + return nil, err + } + fd, err := os.Open(mshellPath) + if err != nil { + return nil, fmt.Errorf("cannot open mshell binary %q: %v", mshellPath, err) + } + return fd, nil +} + +func createPromptAuthKeyFile(fileName string) (string, error) { + fd, err := os.OpenFile(fileName, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0600) + if err != nil { + return "", err + } + defer fd.Close() + keyStr := GenPromptUUID() + _, err = fd.Write([]byte(keyStr)) + if err != nil { + return "", err + } + return keyStr, nil +} + +func ReadPromptAuthKey() (string, error) { + homeDir := GetPromptHomeDir() + err := ensureDir(homeDir) + if err != nil { + return "", fmt.Errorf("cannot find/create PROMPT_HOME directory %q", homeDir) + } + fileName := path.Join(homeDir, PromptAuthKeyFileName) + fd, err := os.Open(fileName) + if err != nil && errors.Is(err, fs.ErrNotExist) { + return createPromptAuthKeyFile(fileName) + } + if err != nil { + return "", fmt.Errorf("error opening prompt authkey:%s: %v", fileName, err) + } + defer fd.Close() + buf, err := io.ReadAll(fd) + if err != nil { + return "", fmt.Errorf("error reading prompt authkey:%s: %v", fileName, err) + } + keyStr := string(buf) + _, err = uuid.Parse(keyStr) + if err != nil { + return "", fmt.Errorf("invalid authkey:%s format: %v", fileName, err) + } + return keyStr, nil +} + +func AcquirePromptLock() (*os.File, error) { + homeDir := GetPromptHomeDir() + err := ensureDir(homeDir) + if err != nil { + return nil, fmt.Errorf("cannot find/create PROMPT_HOME directory %q", homeDir) + } + lockFileName := path.Join(homeDir, PromptLockFile) + fd, err := os.Create(lockFileName) + if err != nil { + return nil, err + } + err = unix.Flock(int(fd.Fd()), unix.LOCK_EX|unix.LOCK_NB) + if err != nil { + fd.Close() + return nil, err + } + return fd, nil +} + +// deprecated (v0.1.8) +func EnsureSessionDir(sessionId string) (string, error) { + if sessionId == "" { + return "", fmt.Errorf("cannot get session dir for blank sessionid") + } + BaseLock.Lock() + sdir, ok := SessionDirCache[sessionId] + BaseLock.Unlock() + if ok { + return sdir, nil + } + scHome := GetPromptHomeDir() + sdir = path.Join(scHome, SessionsDirBaseName, sessionId) + err := ensureDir(sdir) + if err != nil { + return "", err + } + BaseLock.Lock() + SessionDirCache[sessionId] = sdir + BaseLock.Unlock() + return sdir, nil +} + +// deprecated (v0.1.8) +func GetSessionsDir() string { + promptHome := GetPromptHomeDir() + sdir := path.Join(promptHome, SessionsDirBaseName) + return sdir +} + +func EnsureScreenDir(screenId string) (string, error) { + if screenId == "" { + return "", fmt.Errorf("cannot get screen dir for blank sessionid") + } + BaseLock.Lock() + sdir, ok := ScreenDirCache[screenId] + BaseLock.Unlock() + if ok { + return sdir, nil + } + scHome := GetPromptHomeDir() + sdir = path.Join(scHome, ScreensDirBaseName, screenId) + err := ensureDir(sdir) + if err != nil { + return "", err + } + BaseLock.Lock() + ScreenDirCache[screenId] = sdir + BaseLock.Unlock() + return sdir, nil +} + +func GetScreensDir() string { + promptHome := GetPromptHomeDir() + sdir := path.Join(promptHome, ScreensDirBaseName) + return sdir +} + +func ensureDir(dirName string) error { + info, err := os.Stat(dirName) + if errors.Is(err, fs.ErrNotExist) { + err = os.MkdirAll(dirName, 0700) + if err != nil { + return err + } + log.Printf("[prompt] created directory %q\n", dirName) + info, err = os.Stat(dirName) + } + if err != nil { + return err + } + if !info.IsDir() { + return fmt.Errorf("'%s' must be a directory", dirName) + } + return nil +} + +// deprecated (v0.1.8) +func PtyOutFile_Sessions(sessionId string, cmdId string) (string, error) { + sdir, err := EnsureSessionDir(sessionId) + if err != nil { + return "", err + } + if sessionId == "" { + return "", fmt.Errorf("cannot get ptyout file for blank sessionid") + } + if cmdId == "" { + return "", fmt.Errorf("cannot get ptyout file for blank cmdid") + } + return fmt.Sprintf("%s/%s.ptyout.cf", sdir, cmdId), nil +} + +func PtyOutFile(screenId string, lineId string) (string, error) { + sdir, err := EnsureScreenDir(screenId) + if err != nil { + return "", err + } + if screenId == "" { + return "", fmt.Errorf("cannot get ptyout file for blank screenid") + } + if lineId == "" { + return "", fmt.Errorf("cannot get ptyout file for blank lineid") + } + return fmt.Sprintf("%s/%s.ptyout.cf", sdir, lineId), nil +} + +func GenPromptUUID() string { + for { + rtn := uuid.New().String() + _, err := strconv.Atoi(rtn[0:8]) + if err == nil { // do not allow UUIDs where the initial 8 bytes parse to an integer + continue + } + return rtn + } +} + +func NumFormatDec(num int64) string { + var signStr string + absNum := num + if absNum < 0 { + absNum = -absNum + signStr = "-" + } + if absNum < 1000 { + // raw num + return signStr + strconv.FormatInt(absNum, 10) + } + if absNum < 1000000 { + // k num + kVal := float64(absNum) / 1000 + return signStr + strconv.FormatFloat(kVal, 'f', 2, 64) + "k" + } + if absNum < 1000000000 { + // M num + mVal := float64(absNum) / 1000000 + return signStr + strconv.FormatFloat(mVal, 'f', 2, 64) + "m" + } else { + // G num + gVal := float64(absNum) / 1000000000 + return signStr + strconv.FormatFloat(gVal, 'f', 2, 64) + "g" + } +} + +func NumFormatB2(num int64) string { + var signStr string + absNum := num + if absNum < 0 { + absNum = -absNum + signStr = "-" + } + if absNum < 1024 { + // raw num + return signStr + strconv.FormatInt(absNum, 10) + } + if absNum < 1000000 { + // k num + if absNum%1024 == 0 { + return signStr + strconv.FormatInt(absNum/1024, 10) + "K" + } + kVal := float64(absNum) / 1024 + return signStr + strconv.FormatFloat(kVal, 'f', 2, 64) + "K" + } + if absNum < 1000000000 { + // M num + if absNum%(1024*1024) == 0 { + return signStr + strconv.FormatInt(absNum/(1024*1024), 10) + "M" + } + mVal := float64(absNum) / (1024 * 1024) + return signStr + strconv.FormatFloat(mVal, 'f', 2, 64) + "M" + } else { + // G num + if absNum%(1024*1024*1024) == 0 { + return signStr + strconv.FormatInt(absNum/(1024*1024*1024), 10) + "G" + } + gVal := float64(absNum) / (1024 * 1024 * 1024) + return signStr + strconv.FormatFloat(gVal, 'f', 2, 64) + "G" + } +} + +func ClientArch() string { + return fmt.Sprintf("%s/%s", runtime.GOOS, runtime.GOARCH) +} + +var releaseRegex = regexp.MustCompile(`^\d+\.\d+\.\d+$`) +var osReleaseOnce = &sync.Once{} +var osRelease string + +func macOSRelease() string { + ctx, cancelFn := context.WithTimeout(context.Background(), 2*time.Second) + defer cancelFn() + out, err := exec.CommandContext(ctx, "uname", "-r").CombinedOutput() + if err != nil { + log.Printf("error executing uname -r: %v\n", err) + return "-" + } + releaseStr := strings.TrimSpace(string(out)) + if !releaseRegex.MatchString(releaseStr) { + log.Printf("invalid uname -r output: [%s]\n", releaseStr) + return "-" + } + return releaseStr +} + +func MacOSRelease() string { + osReleaseOnce.Do(func() { + osRelease = macOSRelease() + }) + return osRelease +} + +var userShellRegexp = regexp.MustCompile(`^UserShell: (.*)$`) + +// dscl . -read /User/[username] UserShell +// defaults to /bin/bash +func MacUserShell() string { + osUser, err := user.Current() + if err != nil { + log.Printf("error getting current user: %v\n", err) + return DefaultMacOSShell + } + ctx, cancelFn := context.WithTimeout(context.Background(), 2*time.Second) + defer cancelFn() + userStr := "/Users/" + osUser.Name + out, err := exec.CommandContext(ctx, "dscl", ".", "-read", userStr, "UserShell").CombinedOutput() + if err != nil { + log.Printf("error executing macos user shell lookup: %v %q\n", err, string(out)) + return DefaultMacOSShell + } + outStr := strings.TrimSpace(string(out)) + m := userShellRegexp.FindStringSubmatch(outStr) + if m == nil { + log.Printf("error in format of dscl output: %q\n", outStr) + return DefaultMacOSShell + } + return m[1] +} diff --git a/wavesrv/pkg/scpacket/scpacket.go b/wavesrv/pkg/scpacket/scpacket.go new file mode 100644 index 00000000..88255b9e --- /dev/null +++ b/wavesrv/pkg/scpacket/scpacket.go @@ -0,0 +1,120 @@ +package scpacket + +import ( + "fmt" + "reflect" + "strings" + + "github.com/alessio/shellescape" + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/prompt-server/pkg/sstore" +) + +const FeCommandPacketStr = "fecmd" +const WatchScreenPacketStr = "watchscreen" +const FeInputPacketStr = "feinput" +const RemoteInputPacketStr = "remoteinput" + +type FeCommandPacketType struct { + Type string `json:"type"` + MetaCmd string `json:"metacmd"` + MetaSubCmd string `json:"metasubcmd,omitempty"` + Args []string `json:"args,omitempty"` + Kwargs map[string]string `json:"kwargs,omitempty"` + RawStr string `json:"rawstr,omitempty"` + UIContext *UIContextType `json:"uicontext,omitempty"` + Interactive bool `json:"interactive"` +} + +func (pk *FeCommandPacketType) GetRawStr() string { + if pk.RawStr != "" { + return pk.RawStr + } + cmd := "/" + pk.MetaCmd + if pk.MetaSubCmd != "" { + cmd = cmd + ":" + pk.MetaSubCmd + } + var args []string + for k, v := range pk.Kwargs { + argStr := fmt.Sprintf("%s=%s", shellescape.Quote(k), shellescape.Quote(v)) + args = append(args, argStr) + } + for _, arg := range pk.Args { + args = append(args, shellescape.Quote(arg)) + } + if len(args) == 0 { + return cmd + } + return cmd + " " + strings.Join(args, " ") +} + +type UIContextType struct { + SessionId string `json:"sessionid"` + ScreenId string `json:"screenid"` + Remote *sstore.RemotePtrType `json:"remote,omitempty"` + WinSize *packet.WinSize `json:"winsize,omitempty"` + Build string `json:"build,omitempty"` +} + +type FeInputPacketType struct { + Type string `json:"type"` + CK base.CommandKey `json:"ck"` + Remote sstore.RemotePtrType `json:"remote"` + InputData64 string `json:"inputdata64"` + SigName string `json:"signame,omitempty"` + WinSize *packet.WinSize `json:"winsize,omitempty"` +} + +type RemoteInputPacketType struct { + Type string `json:"type"` + RemoteId string `json:"remoteid"` + InputData64 string `json:"inputdata64"` +} + +type WatchScreenPacketType struct { + Type string `json:"type"` + SessionId string `json:"sessionid"` + ScreenId string `json:"screenid"` + Connect bool `json:"connect"` + AuthKey string `json:"authkey"` +} + +func init() { + packet.RegisterPacketType(FeCommandPacketStr, reflect.TypeOf(FeCommandPacketType{})) + packet.RegisterPacketType(WatchScreenPacketStr, reflect.TypeOf(WatchScreenPacketType{})) + packet.RegisterPacketType(FeInputPacketStr, reflect.TypeOf(FeInputPacketType{})) + packet.RegisterPacketType(RemoteInputPacketStr, reflect.TypeOf(RemoteInputPacketType{})) +} + +func (*FeCommandPacketType) GetType() string { + return FeCommandPacketStr +} + +func MakeFeCommandPacket() *FeCommandPacketType { + return &FeCommandPacketType{Type: FeCommandPacketStr} +} + +func (*FeInputPacketType) GetType() string { + return FeInputPacketStr +} + +func MakeFeInputPacket() *FeInputPacketType { + return &FeInputPacketType{Type: FeInputPacketStr} +} + +func (*WatchScreenPacketType) GetType() string { + return WatchScreenPacketStr +} + +func MakeWatchScreenPacket() *WatchScreenPacketType { + return &WatchScreenPacketType{Type: WatchScreenPacketStr} +} + +func MakeRemoteInputPacket() *RemoteInputPacketType { + return &RemoteInputPacketType{Type: RemoteInputPacketStr} +} + +func (*RemoteInputPacketType) GetType() string { + return RemoteInputPacketStr +} diff --git a/wavesrv/pkg/scws/scws.go b/wavesrv/pkg/scws/scws.go new file mode 100644 index 00000000..57b1c65d --- /dev/null +++ b/wavesrv/pkg/scws/scws.go @@ -0,0 +1,305 @@ +package scws + +import ( + "context" + "fmt" + "log" + "sync" + "time" + + "github.com/google/uuid" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/prompt-server/pkg/mapqueue" + "github.com/commandlinedev/prompt-server/pkg/remote" + "github.com/commandlinedev/prompt-server/pkg/scpacket" + "github.com/commandlinedev/prompt-server/pkg/sstore" + "github.com/commandlinedev/prompt-server/pkg/wsshell" +) + +const WSStatePacketChSize = 20 +const MaxInputDataSize = 1000 +const RemoteInputQueueSize = 100 + +var RemoteInputMapQueue *mapqueue.MapQueue + +func init() { + RemoteInputMapQueue = mapqueue.MakeMapQueue(RemoteInputQueueSize) +} + +type WSState struct { + Lock *sync.Mutex + ClientId string + ConnectTime time.Time + Shell *wsshell.WSShell + UpdateCh chan interface{} + UpdateQueue []interface{} + Authenticated bool + AuthKey string + + SessionId string + ScreenId string +} + +func MakeWSState(clientId string, authKey string) *WSState { + rtn := &WSState{} + rtn.Lock = &sync.Mutex{} + rtn.ClientId = clientId + rtn.ConnectTime = time.Now() + rtn.AuthKey = authKey + return rtn +} + +func (ws *WSState) SetAuthenticated(authVal bool) { + ws.Lock.Lock() + defer ws.Lock.Unlock() + ws.Authenticated = authVal +} + +func (ws *WSState) IsAuthenticated() bool { + ws.Lock.Lock() + defer ws.Lock.Unlock() + return ws.Authenticated +} + +func (ws *WSState) GetShell() *wsshell.WSShell { + ws.Lock.Lock() + defer ws.Lock.Unlock() + return ws.Shell +} + +func (ws *WSState) WriteUpdate(update interface{}) error { + shell := ws.GetShell() + if shell == nil { + return fmt.Errorf("cannot write update, empty shell") + } + err := shell.WriteJson(update) + if err != nil { + return err + } + return nil +} + +func (ws *WSState) UpdateConnectTime() { + ws.Lock.Lock() + defer ws.Lock.Unlock() + ws.ConnectTime = time.Now() +} + +func (ws *WSState) GetConnectTime() time.Time { + ws.Lock.Lock() + defer ws.Lock.Unlock() + return ws.ConnectTime +} + +func (ws *WSState) WatchScreen(sessionId string, screenId string) { + ws.Lock.Lock() + defer ws.Lock.Unlock() + if ws.SessionId == sessionId && ws.ScreenId == screenId { + return + } + ws.SessionId = sessionId + ws.ScreenId = screenId + ws.UpdateCh = sstore.MainBus.RegisterChannel(ws.ClientId, ws.ScreenId) + go ws.RunUpdates(ws.UpdateCh) +} + +func (ws *WSState) UnWatchScreen() { + ws.Lock.Lock() + defer ws.Lock.Unlock() + sstore.MainBus.UnregisterChannel(ws.ClientId) + ws.SessionId = "" + ws.ScreenId = "" + log.Printf("[ws] unwatch screen clientid=%s\n", ws.ClientId) +} + +func (ws *WSState) getUpdateCh() chan interface{} { + ws.Lock.Lock() + defer ws.Lock.Unlock() + return ws.UpdateCh +} + +func (ws *WSState) RunUpdates(updateCh chan interface{}) { + if updateCh == nil { + panic("invalid nil updateCh passed to RunUpdates") + } + for update := range updateCh { + shell := ws.GetShell() + if shell != nil { + shell.WriteJson(update) + } + } +} + +func (ws *WSState) ReplaceShell(shell *wsshell.WSShell) { + ws.Lock.Lock() + defer ws.Lock.Unlock() + if ws.Shell == nil { + ws.Shell = shell + return + } + ws.Shell.Conn.Close() + ws.Shell = shell + return +} + +func (ws *WSState) handleConnection() error { + ctx, cancelFn := context.WithTimeout(context.Background(), 5*time.Second) + defer cancelFn() + update, err := sstore.GetAllSessions(ctx) + if err != nil { + return fmt.Errorf("getting sessions: %w", err) + } + remotes := remote.GetAllRemoteRuntimeState() + ifarr := make([]interface{}, len(remotes)) + for idx, r := range remotes { + ifarr[idx] = r + } + update.Remotes = ifarr + update.Connect = true + err = ws.Shell.WriteJson(update) + if err != nil { + return err + } + return nil +} + +func (ws *WSState) handleWatchScreen(wsPk *scpacket.WatchScreenPacketType) error { + if wsPk.SessionId != "" { + if _, err := uuid.Parse(wsPk.SessionId); err != nil { + return fmt.Errorf("invalid watchscreen sessionid: %w", err) + } + } + if wsPk.ScreenId != "" { + if _, err := uuid.Parse(wsPk.ScreenId); err != nil { + return fmt.Errorf("invalid watchscreen screenid: %w", err) + } + } + if wsPk.AuthKey == "" { + ws.SetAuthenticated(false) + return fmt.Errorf("invalid watchscreen, no authkey") + } + if wsPk.AuthKey != ws.AuthKey { + ws.SetAuthenticated(false) + return fmt.Errorf("invalid watchscreen, invalid authkey") + } + ws.SetAuthenticated(true) + if wsPk.SessionId == "" || wsPk.ScreenId == "" { + ws.UnWatchScreen() + } else { + ws.WatchScreen(wsPk.SessionId, wsPk.ScreenId) + log.Printf("[ws %s] watchscreen %s/%s\n", ws.ClientId, wsPk.SessionId, wsPk.ScreenId) + } + if wsPk.Connect { + // log.Printf("[ws %s] watchscreen connect\n", ws.ClientId) + err := ws.handleConnection() + if err != nil { + return fmt.Errorf("connect: %w", err) + } + } + return nil +} + +func (ws *WSState) RunWSRead() { + shell := ws.GetShell() + if shell == nil { + return + } + shell.WriteJson(map[string]interface{}{"type": "hello"}) // let client know we accepted this connection, ignore error + for msgBytes := range shell.ReadChan { + pk, err := packet.ParseJsonPacket(msgBytes) + if err != nil { + log.Printf("error unmarshalling ws message: %v\n", err) + continue + } + if pk.GetType() == scpacket.WatchScreenPacketStr { + wsPk := pk.(*scpacket.WatchScreenPacketType) + err := ws.handleWatchScreen(wsPk) + if err != nil { + // TODO send errors back to client, likely unrecoverable + log.Printf("[ws %s] error %v\n", ws.ClientId, err) + } + continue + } + isAuth := ws.IsAuthenticated() + if !isAuth { + log.Printf("[error] cannot process ws-packet[%s], not authenticated\n", pk.GetType()) + continue + } + if pk.GetType() == scpacket.FeInputPacketStr { + feInputPk := pk.(*scpacket.FeInputPacketType) + if feInputPk.Remote.OwnerId != "" { + log.Printf("[error] cannot send input to remote with ownerid\n") + continue + } + if feInputPk.Remote.RemoteId == "" { + log.Printf("[error] invalid input packet, remoteid is not set\n") + continue + } + err := RemoteInputMapQueue.Enqueue(feInputPk.Remote.RemoteId, func() { + err = sendCmdInput(feInputPk) + if err != nil { + log.Printf("[error] sending command input: %v\n", err) + } + }) + if err != nil { + log.Printf("[error] could not queue sendCmdInput: %v\n", err) + continue + } + continue + } + if pk.GetType() == scpacket.RemoteInputPacketStr { + inputPk := pk.(*scpacket.RemoteInputPacketType) + if inputPk.RemoteId == "" { + log.Printf("[error] invalid remoteinput packet, remoteid is not set\n") + continue + } + go func() { + err = remote.SendRemoteInput(inputPk) + if err != nil { + log.Printf("[error] processing remote input: %v\n", err) + } + }() + continue + } + log.Printf("got ws bad message: %v\n", pk.GetType()) + } +} + +func sendCmdInput(pk *scpacket.FeInputPacketType) error { + err := pk.CK.Validate("input packet") + if err != nil { + return err + } + if pk.Remote.RemoteId == "" { + return fmt.Errorf("input must set remoteid") + } + msh := remote.GetRemoteById(pk.Remote.RemoteId) + if msh == nil { + return fmt.Errorf("remote %s not found", pk.Remote.RemoteId) + } + if len(pk.InputData64) > 0 { + inputLen := packet.B64DecodedLen(pk.InputData64) + if inputLen > MaxInputDataSize { + return fmt.Errorf("input data size too large, len=%d (max=%d)", inputLen, MaxInputDataSize) + } + dataPk := packet.MakeDataPacket() + dataPk.CK = pk.CK + dataPk.FdNum = 0 // stdin + dataPk.Data64 = pk.InputData64 + err = msh.SendInput(dataPk) + if err != nil { + return err + } + } + if pk.SigName != "" || pk.WinSize != nil { + siPk := packet.MakeSpecialInputPacket() + siPk.CK = pk.CK + siPk.SigName = pk.SigName + siPk.WinSize = pk.WinSize + err = msh.SendSpecialInput(siPk) + if err != nil { + return err + } + } + return nil +} diff --git a/wavesrv/pkg/shparse/comp.go b/wavesrv/pkg/shparse/comp.go new file mode 100644 index 00000000..28ddde7b --- /dev/null +++ b/wavesrv/pkg/shparse/comp.go @@ -0,0 +1,288 @@ +package shparse + +import ( + "strings" + + "github.com/commandlinedev/prompt-server/pkg/utilfn" +) + +const ( + CompTypeCommandMeta = "command-meta" + CompTypeCommand = "command" + CompTypeArg = "command-arg" + CompTypeInvalid = "invalid" + CompTypeVar = "var" + CompTypeAssignment = "assignment" + CompTypeBasic = "basic" +) + +type CompletionPos struct { + RawPos int // the raw position of cursor + SuperOffset int // adjust all offsets in Cmd and CmdWord by SuperOffset + + CompType string // see CompType* constants + Cmd *CmdType // nil if between commands or a special completion (otherwise will be a SimpleCommand) + // index into cmd.Words (only set when Cmd is not nil, otherwise we look at CompCommand) + // 0 means command-word + // negative means assignment-words. + // can be past the end of Words (means start new word). + CmdWordPos int + CompWord *WordType // set to the word we are completing (nil if we are starting a new word) + CompWordOffset int // offset into compword (only if CmdWord is not nil) +} + +func compTypeFromPos(cmdWordPos int) string { + if cmdWordPos == 0 { + return CompTypeCommand + } + if cmdWordPos < 0 { + return CompTypeAssignment + } + return CompTypeArg +} + +func (cmd *CmdType) findCompletionPos_simple(pos int, superOffset int) CompletionPos { + if cmd.Type != CmdTypeSimple { + panic("findCompletetionPos_simple only works for CmdTypeSimple") + } + rtn := CompletionPos{RawPos: pos, SuperOffset: superOffset, Cmd: cmd} + for idx, word := range cmd.AssignmentWords { + startOffset := word.Offset + endOffset := word.Offset + len(word.Raw) + if pos <= startOffset { + // starting a new word at this position (before the current assignment word) + rtn.CmdWordPos = idx - len(cmd.AssignmentWords) + rtn.CompType = CompTypeAssignment + return rtn + } + if pos <= endOffset { + // completing an assignment word + rtn.CmdWordPos = idx - len(cmd.AssignmentWords) + rtn.CompWord = word + rtn.CompWordOffset = pos - word.Offset + rtn.CompType = CompTypeAssignment + return rtn + } + } + var foundWord *WordType + var foundWordIdx int + for idx, word := range cmd.Words { + startOffset := word.Offset + endOffset := word.Offset + len(word.Raw) + if pos <= startOffset { + // starting a new word at this position + rtn.CmdWordPos = idx + rtn.CompType = compTypeFromPos(idx) + return rtn + } + if pos == endOffset && word.Type == WordTypeOp { + // operators are special, they can allow a full-word completion at endpos + continue + } + if pos <= endOffset { + foundWord = word + foundWordIdx = idx + break + } + } + if foundWord != nil { + rtn.CmdWordPos = foundWordIdx + rtn.CompWord = foundWord + rtn.CompWordOffset = pos - foundWord.Offset + if foundWord.uncompletable() { + // invalid completion point + rtn.CompType = CompTypeInvalid + return rtn + } + rtn.CompType = compTypeFromPos(foundWordIdx) + return rtn + } + // past the end, so we're starting a new word in Cmd + rtn.CmdWordPos = len(cmd.Words) + rtn.CompType = CompTypeArg + return rtn +} + +func (cmd *CmdType) findCompletionPos_none(pos int, superOffset int) CompletionPos { + rtn := CompletionPos{RawPos: pos, SuperOffset: superOffset} + if cmd.Type != CmdTypeNone { + panic("findCompletionPos_none only works for CmdTypeNone") + } + var foundWord *WordType + for _, word := range cmd.Words { + startOffset := word.Offset + endOffset := word.Offset + len(word.Raw) + if pos <= startOffset { + break + } + if pos <= endOffset { + if pos == endOffset && word.Type == WordTypeOp { + // operators are special, they can allow a full-word completion at endpos + continue + } + foundWord = word + break + } + } + if foundWord == nil { + // just revert to a file completion + rtn.CompType = CompTypeBasic + return rtn + } + foundWordOffset := pos - foundWord.Offset + rtn.CompWord = foundWord + rtn.CompWordOffset = foundWordOffset + if foundWord.uncompletable() { + // ok, we're inside of a word in CmdTypeNone. if we're in an uncompletable word, return CompInvalid + rtn.CompType = CompTypeInvalid + return rtn + } + if foundWordOffset > 0 && foundWordOffset < foundWord.contentStartPos() { + // cursor is in a weird position, between characters of a multi-char prefix (e.g. "$[*]{hello}" or $[*]'hello'). cannot complete. + rtn.CompType = CompTypeInvalid + return rtn + } + // revert to file completion + rtn.CompType = CompTypeBasic + return rtn +} + +func findCompletionWordAtPos(words []*WordType, pos int, allowEndMatch bool) *WordType { + // WordTypeSimpleVar is special (always allowEndMatch), if cursor is at the end of SimpleVar it is returned + for _, word := range words { + if pos > word.Offset && pos < word.Offset+len(word.Raw) { + return word + } + if (allowEndMatch || word.Type == WordTypeSimpleVar) && pos == word.Offset+len(word.Raw) { + return word + } + } + return nil +} + +// recursively descend down the word, parse commands and find a sub completion point if any. +// return nil if there is no sub completion point in this word +func findCompletionPosInWord(word *WordType, posInWord int, superOffset int) *CompletionPos { + rawPos := word.Offset + posInWord + if word.Type == WordTypeGroup || word.Type == WordTypeDQ || word.Type == WordTypeDDQ { + // need to descend further + if posInWord <= word.contentStartPos() { + return nil + } + if posInWord > word.contentEndPos() { + return nil + } + subWord := findCompletionWordAtPos(word.Subs, posInWord-word.contentStartPos(), false) + if subWord == nil { + return nil + } + return findCompletionPosInWord(subWord, posInWord-(subWord.Offset+word.contentStartPos()), superOffset+(word.Offset+word.contentStartPos())) + } + if word.Type == WordTypeDP || word.Type == WordTypeBQ { + if posInWord < word.contentStartPos() { + return nil + } + if posInWord > word.contentEndPos() { + return nil + } + subCmds := ParseCommands(word.Subs) + newPos := findCompletionPosInternal(subCmds, posInWord-word.contentStartPos(), superOffset+(word.Offset+word.contentStartPos())) + return &newPos + } + if word.Type == WordTypeSimpleVar || word.Type == WordTypeVarBrace { + // special "var" completion + rtn := &CompletionPos{RawPos: rawPos, SuperOffset: superOffset} + rtn.CompType = CompTypeVar + rtn.CompWordOffset = posInWord + rtn.CompWord = word + return rtn + } + return nil +} + +// returns the context for completion +// if we are completing in a simple-command, the returns the Cmd. the Cmd can be used for specialized completion (command name, arg position, etc.) +// if we are completing in a word, returns the Word. Word might be a group-word or DQ word, so it may need additional resolution (done in extend) +// otherwise we are going to create a new word to insert at offset (so the context does not matter) +func findCompletionPosCmds(cmds []*CmdType, pos int, superOffset int) CompletionPos { + rtn := CompletionPos{RawPos: pos, SuperOffset: superOffset} + if len(cmds) == 0 { + // set CompCommand because we're starting a new command + rtn.CompType = CompTypeCommand + return rtn + } + for _, cmd := range cmds { + endOffset := cmd.endOffset() + if pos > endOffset || (cmd.Type == CmdTypeNone && pos == endOffset) { + continue + } + startOffset := cmd.offset() + if cmd.Type == CmdTypeSimple { + if pos <= startOffset { + rtn.CompType = CompTypeCommand + return rtn + } + return cmd.findCompletionPos_simple(pos, superOffset) + } else { + // not in a simple-command + // if we're before the none-command, just start a new command + if pos <= startOffset { + rtn.CompType = CompTypeCommand + return rtn + } + return cmd.findCompletionPos_none(pos, superOffset) + } + } + // past the end + lastCmd := cmds[len(cmds)-1] + if lastCmd.Type == CmdTypeSimple { + // just extend last command + rtn.Cmd = lastCmd + rtn.CmdWordPos = len(lastCmd.Words) + rtn.CompType = CompTypeArg + return rtn + } + // use lastCmd.NoneComplete to see if last command ended on a "separator". use that to set CompCommand + if lastCmd.NoneComplete { + rtn.CompType = CompTypeCommand + } else { + rtn.CompType = CompTypeBasic + } + return rtn +} + +func findCompletionPosInternal(cmds []*CmdType, pos int, superOffset int) CompletionPos { + cpos := findCompletionPosCmds(cmds, pos, superOffset) + if cpos.CompWord == nil { + return cpos + } + subPos := findCompletionPosInWord(cpos.CompWord, cpos.CompWordOffset, superOffset) + if subPos != nil { + return *subPos + } + return cpos +} + +func FindCompletionPos(cmds []*CmdType, pos int) CompletionPos { + cpos := findCompletionPosInternal(cmds, pos, 0) + if cpos.CompType == CompTypeCommand && cpos.SuperOffset == 0 && cpos.CompWord != nil && cpos.CompWord.Offset == 0 && strings.HasPrefix(string(cpos.CompWord.Raw), "/") { + cpos.CompType = CompTypeCommandMeta + } + return cpos +} + +func (cpos CompletionPos) Extend(origStr utilfn.StrWithPos, extensionStr string, extensionComplete bool) utilfn.StrWithPos { + compWord := cpos.CompWord + if compWord == nil { + compWord = MakeEmptyWord(WordTypeLit, nil, cpos.RawPos, true) + } + realOffset := compWord.Offset + cpos.SuperOffset + if strings.HasSuffix(extensionStr, "/") { + extensionComplete = false + } + rtnSP := Extend(compWord, cpos.CompWordOffset, extensionStr, extensionComplete) + origRunes := []rune(origStr.Str) + rtnSP = rtnSP.Prepend(string(origRunes[0:realOffset])) + rtnSP = rtnSP.Append(string(origRunes[realOffset+len(compWord.Raw):])) + return rtnSP +} diff --git a/wavesrv/pkg/shparse/expand.go b/wavesrv/pkg/shparse/expand.go new file mode 100644 index 00000000..9f872311 --- /dev/null +++ b/wavesrv/pkg/shparse/expand.go @@ -0,0 +1,258 @@ +package shparse + +import ( + "bytes" + "fmt" + + "mvdan.cc/sh/v3/expand" +) + +const MaxExpandLen = 64 * 1024 + +type ExpandInfo struct { + HasTilde bool // only ~ as the first character when SimpleExpandContext.HomeDir is set + HasVar bool // $x, $$, ${...} + HasGlob bool // *, ?, [, { + HasExtGlob bool // ?(...) ... ?*+@! + HasHistory bool // ! (anywhere) + HasSpecial bool // subshell, arith +} + +type ExpandContext struct { + HomeDir string +} + +func expandSQ(buf *bytes.Buffer, rawLit []rune) { + // no info specials + buf.WriteString(string(rawLit)) +} + +// TODO implement our own ANSI single quote formatter +func expandANSISQ(buf *bytes.Buffer, rawLit []rune) { + // no info specials + str, _, _ := expand.Format(nil, string(rawLit), nil) + buf.WriteString(str) +} + +func expandLiteral(buf *bytes.Buffer, info *ExpandInfo, rawLit []rune) { + var lastBackSlash bool + var lastExtGlob bool + var lastDollar bool + for _, ch := range rawLit { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '\n' { + // special case, backslash *and* newline are ignored + continue + } + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + lastExtGlob = false + lastDollar = false + continue + } + if ch == '*' || ch == '?' || ch == '[' || ch == '{' { + info.HasGlob = true + } + if ch == '`' { + info.HasSpecial = true + } + if ch == '!' { + info.HasHistory = true + } + if lastExtGlob && ch == '(' { + info.HasExtGlob = true + } + if lastDollar && (ch != ' ' && ch != '"' && ch != '\'' && ch != '(' || ch != '[') { + info.HasVar = true + } + if lastDollar && (ch == '(' || ch == '[') { + info.HasSpecial = true + } + lastExtGlob = (ch == '?' || ch == '*' || ch == '+' || ch == '@' || ch == '!') + lastDollar = (ch == '$') + buf.WriteRune(ch) + } + if lastBackSlash { + buf.WriteByte('\\') + } +} + +// will also work for partial double quoted strings +func expandDQLiteral(buf *bytes.Buffer, info *ExpandInfo, rawVal []rune) { + var lastBackSlash bool + var lastDollar bool + for _, ch := range rawVal { + if ch == 0 { + break + } + if lastBackSlash { + lastBackSlash = false + if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { + buf.WriteRune(ch) + continue + } + buf.WriteRune('\\') + buf.WriteRune(ch) + continue + } + if ch == '\\' { + lastBackSlash = true + lastDollar = false + continue + } + + // similar to expandLiteral, but no globbing + if ch == '`' { + info.HasSpecial = true + } + if ch == '!' { + info.HasHistory = true + } + if lastDollar && (ch != ' ' && ch != '"' && ch != '\'' && ch != '(' || ch != '[') { + info.HasVar = true + } + if lastDollar && (ch == '(' || ch == '[') { + info.HasSpecial = true + } + lastDollar = (ch == '$') + buf.WriteRune(ch) + } + // in a valid parsed DQ string, you cannot have a trailing backslash (because \" would not end the string) + // still putting the case here though in case we ever deal with incomplete strings (e.g. completion) + if lastBackSlash { + buf.WriteByte('\\') + } +} + +func simpleExpandSubs(buf *bytes.Buffer, info *ExpandInfo, ectx ExpandContext, word *WordType, pos int) { + fmt.Printf("expand subs: %v\n", word) + parts := word.Subs + startPos := word.contentStartPos() + for _, part := range parts { + remainingLen := pos - startPos + if remainingLen <= 0 { + break + } + simpleExpandWord(buf, info, ectx, part, remainingLen) + startPos += len(part.Raw) + } +} + +func canExpand(ectx ExpandContext, wtype string) bool { + return wtype == WordTypeLit || wtype == WordTypeSQ || wtype == WordTypeDSQ || + wtype == WordTypeDQ || wtype == WordTypeDDQ || wtype == WordTypeGroup +} + +func simpleExpandWord(buf *bytes.Buffer, info *ExpandInfo, ectx ExpandContext, word *WordType, pos int) { + if canExpand(ectx, word.Type) { + if pos >= word.contentEndPos() { + pos = word.contentEndPos() + } + if pos <= word.contentStartPos() { + return + } + } else { + if pos >= len(word.Raw) { + pos = len(word.Raw) + } + if pos <= 0 { + return + } + } + + switch word.Type { + case WordTypeLit: + if word.QC.cur() == WordTypeDQ { + expandDQLiteral(buf, info, word.Raw[:pos]) + return + } + expandLiteral(buf, info, word.Raw[:pos]) + + case WordTypeSQ: + expandSQ(buf, word.Raw[word.contentStartPos():pos]) + return + + case WordTypeDSQ: + expandANSISQ(buf, word.Raw[word.contentStartPos():pos]) + return + + case WordTypeDQ, WordTypeDDQ: + simpleExpandSubs(buf, info, ectx, word, pos) + return + + case WordTypeGroup: + simpleExpandSubs(buf, info, ectx, word, pos) + return + + // not expanded + case WordTypeSimpleVar: + info.HasVar = true + buf.WriteString(string(word.Raw[:pos])) + return + + // not expanded + case WordTypeVarBrace: + info.HasVar = true + buf.WriteString(string(word.Raw[:pos])) + return + + default: + info.HasSpecial = true + buf.WriteString(string(word.Raw[:pos])) + return + } +} + +func SimpleExpandPrefix(ectx ExpandContext, word *WordType, pos int) (string, ExpandInfo) { + var buf bytes.Buffer + var info ExpandInfo + simpleExpandWord(&buf, &info, ectx, word, pos) + return buf.String(), info +} + +func SimpleExpand(ectx ExpandContext, word *WordType) (string, ExpandInfo) { + return SimpleExpandPrefix(ectx, word, len(word.Raw)) +} + +// returns varname (no '$') and ok (whether this is a valid varname expansion) +func SimpleVarNamePrefix(ectx ExpandContext, word *WordType, pos int) (string, bool) { + if word.Type != WordTypeSimpleVar && word.Type != WordTypeVarBrace { + return "", false + } + if word.Type == WordTypeSimpleVar { + if pos == 0 { + return "", false + } + if pos == 1 { + return "", true + } + if pos > len(word.Raw) { + pos = len(word.Raw) + } + return string(word.Raw[1:pos]), true + } + + // word.Type == WordTypeVarBrace + // knock '${' off the front, then see if the rest is a valid var name. + if pos == 0 || pos == 1 { + return "", false + } + if pos == 2 { + return "", true + } + if pos > word.contentEndPos() { + pos = word.contentEndPos() + } + rawVarName := word.Raw[2:pos] + if isSimpleVarName(rawVarName) { + return string(rawVarName), true + } + return "", false +} diff --git a/wavesrv/pkg/shparse/extend.go b/wavesrv/pkg/shparse/extend.go new file mode 100644 index 00000000..8da19810 --- /dev/null +++ b/wavesrv/pkg/shparse/extend.go @@ -0,0 +1,410 @@ +package shparse + +import ( + "bytes" + "unicode" + "unicode/utf8" + + "github.com/commandlinedev/prompt-server/pkg/utilfn" +) + +var noEscChars []bool +var specialEsc []string + +func init() { + noEscChars = make([]bool, 256) + for ch := 0; ch < 256; ch++ { + if (ch >= '0' && ch <= '9') || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || + ch == '-' || ch == '.' || ch == '/' || ch == ':' || ch == '=' || ch == '_' { + noEscChars[byte(ch)] = true + } + } + specialEsc = make([]string, 256) + specialEsc[0x7] = "\\a" + specialEsc[0x8] = "\\b" + specialEsc[0x9] = "\\t" + specialEsc[0xa] = "\\n" + specialEsc[0xb] = "\\v" + specialEsc[0xc] = "\\f" + specialEsc[0xd] = "\\r" + specialEsc[0x1b] = "\\E" +} + +func getUtf8Literal(ch rune) string { + var buf bytes.Buffer + var runeArr [utf8.UTFMax]byte + barr := runeArr[:] + byteLen := utf8.EncodeRune(barr, ch) + for i := 0; i < byteLen; i++ { + buf.WriteString("\\x") + buf.WriteByte(utilfn.HexDigits[barr[i]/16]) + buf.WriteByte(utilfn.HexDigits[barr[i]%16]) + } + return buf.String() +} + +func (w *WordType) writeString(s string) { + for _, ch := range s { + w.writeRune(ch) + } +} + +func (w *WordType) writeRune(ch rune) { + wmeta := wordMetaMap[w.Type] + if w.Complete && wmeta.SuffixLen == 1 { + w.Raw = append(w.Raw[0:len(w.Raw)-1], ch, w.Raw[len(w.Raw)-1]) + return + } + if w.Complete && wmeta.SuffixLen == 2 { + w.Raw = append(w.Raw[0:len(w.Raw)-2], ch, w.Raw[len(w.Raw)-2], w.Raw[len(w.Raw)-1]) + return + } + // not complete or SuffixLen == 0 (2+ is not supported) + w.Raw = append(w.Raw, ch) + return +} + +type extendContext struct { + Input []*WordType + InputPos int + QC QuoteContext + Rtn []*WordType + CurWord *WordType + Intention string +} + +func makeExtendContext(qc QuoteContext, word *WordType) *extendContext { + rtn := &extendContext{QC: qc} + if word == nil { + rtn.Intention = WordTypeLit + return rtn + } else { + rtn.Intention = word.Type + rtn.Rtn = []*WordType{word} + rtn.CurWord = word + return rtn + } +} + +func (ec *extendContext) appendWord(w *WordType) { + ec.Rtn = append(ec.Rtn, w) + ec.CurWord = w +} + +func (ec *extendContext) ensureCurWord() { + if ec.CurWord == nil || ec.CurWord.Type != ec.Intention { + ec.CurWord = MakeEmptyWord(ec.Intention, ec.QC, 0, true) + ec.Rtn = append(ec.Rtn, ec.CurWord) + } +} + +// grp, dq, ddq +func extendWithSubs(word *WordType, wordPos int, extStr string, complete bool) utilfn.StrWithPos { + wmeta := wordMetaMap[word.Type] + if word.Type == WordTypeGroup { + atEnd := (wordPos == len(word.Raw)) + subWord := findCompletionWordAtPos(word.Subs, wordPos, true) + if subWord == nil { + strPos := Extend(MakeEmptyWord(WordTypeLit, word.QC, 0, true), 0, extStr, atEnd) + strPos = strPos.Prepend(string(word.Raw[0:wordPos])) + strPos = strPos.Append(string(word.Raw[wordPos:])) + return strPos + } else { + subComplete := complete && atEnd + strPos := Extend(subWord, wordPos-subWord.Offset, extStr, subComplete) + strPos = strPos.Prepend(string(word.Raw[0:subWord.Offset])) + strPos = strPos.Append(string(word.Raw[subWord.Offset+len(subWord.Raw):])) + return strPos + } + } else if word.Type == WordTypeDQ || word.Type == WordTypeDDQ { + if wordPos < word.contentStartPos() { + wordPos = word.contentStartPos() + } + atEnd := (wordPos >= len(word.Raw)-wmeta.SuffixLen) + subWord := findCompletionWordAtPos(word.Subs, wordPos-wmeta.PrefixLen, true) + quoteBalance := !atEnd + if subWord == nil { + realOffset := wordPos + strPos, wordOpen := extendInternal(MakeEmptyWord(WordTypeLit, word.QC.push(WordTypeDQ), 0, true), 0, extStr, false, quoteBalance) + strPos = strPos.Prepend(string(word.Raw[0:realOffset])) + var requiredSuffix string + if wordOpen { + requiredSuffix = wmeta.getSuffix() + } + if atEnd { + if complete { + return utilfn.StrWithPos{Str: strPos.Str + requiredSuffix + " ", Pos: strPos.Pos + len(requiredSuffix) + 1} + } else { + if word.Complete && requiredSuffix != "" { + return strPos.Append(requiredSuffix) + } + return strPos + } + } + strPos = strPos.Append(string(word.Raw[wordPos:])) + return strPos + } else { + realOffset := subWord.Offset + wmeta.PrefixLen + strPos, wordOpen := extendInternal(subWord, wordPos-realOffset, extStr, false, quoteBalance) + strPos = strPos.Prepend(string(word.Raw[0:realOffset])) + var requiredSuffix string + if wordOpen { + requiredSuffix = wmeta.getSuffix() + } + if atEnd { + if complete { + return utilfn.StrWithPos{Str: strPos.Str + requiredSuffix + " ", Pos: strPos.Pos + len(requiredSuffix) + 1} + } else { + if word.Complete && requiredSuffix != "" { + return strPos.Append(requiredSuffix) + } + return strPos + } + } + strPos = strPos.Append(string(word.Raw[realOffset+len(subWord.Raw):])) + return strPos + } + } else { + return utilfn.StrWithPos{Str: string(word.Raw), Pos: wordPos} + } +} + +// lit, svar, varb, sq, dsq +func extendLeafCh(buf *bytes.Buffer, wordOpen *bool, wtype string, qc QuoteContext, ch rune) { + switch wtype { + case WordTypeSimpleVar, WordTypeVarBrace: + extendVar(buf, ch) + + case WordTypeLit: + if qc.cur() == WordTypeDQ { + extendDQLit(buf, wordOpen, ch) + } else { + extendLit(buf, ch) + } + + case WordTypeSQ: + extendSQ(buf, wordOpen, ch) + + case WordTypeDSQ: + extendDSQ(buf, wordOpen, ch) + + default: + return + } +} + +func getWordOpenStr(wtype string, qc QuoteContext) string { + if wtype == WordTypeLit { + if qc.cur() == WordTypeDQ { + return "\"" + } else { + return "" + } + } + wmeta := wordMetaMap[wtype] + return wmeta.getPrefix() +} + +// lit, svar, varb sq, dsq +func extendLeaf(buf *bytes.Buffer, wordOpen *bool, word *WordType, wordPos int, extStr string) { + for _, ch := range extStr { + extendLeafCh(buf, wordOpen, word.Type, word.QC, ch) + } +} + +// lit, grp, svar, dq, ddq, varb, sq, dsq +// returns (strwithpos, dq-closed) +func extendInternal(word *WordType, wordPos int, extStr string, complete bool, requiresQuoteBalance bool) (utilfn.StrWithPos, bool) { + if extStr == "" { + return utilfn.StrWithPos{Str: string(word.Raw), Pos: wordPos}, true + } + if word.canHaveSubs() { + return extendWithSubs(word, wordPos, extStr, complete), true + } + var buf bytes.Buffer + isEOW := wordPos >= word.contentEndPos() + if isEOW { + wordPos = word.contentEndPos() + } + if wordPos < word.contentStartPos() { + wordPos = word.contentStartPos() + } + if wordPos > 0 { + buf.WriteString(string(word.Raw[0:word.contentStartPos()])) // write the prefix + } + if wordPos > word.contentStartPos() { + buf.WriteString(string(word.Raw[word.contentStartPos():wordPos])) + } + wordOpen := true + extendLeaf(&buf, &wordOpen, word, wordPos, extStr) + if isEOW { + // end-of-word, write the suffix (and optional ' '). return the end of the string + wmeta := wordMetaMap[word.Type] + rtnPos := utf8.RuneCount(buf.Bytes()) + buf.WriteString(wmeta.getSuffix()) + if !wordOpen && requiresQuoteBalance { + buf.WriteString(getWordOpenStr(word.Type, word.QC)) + wordOpen = true + } + if complete { + buf.WriteRune(' ') + return utilfn.StrWithPos{Str: buf.String(), Pos: utf8.RuneCount(buf.Bytes())}, wordOpen + } else { + return utilfn.StrWithPos{Str: buf.String(), Pos: rtnPos}, wordOpen + } + } + // completion in the middle of a word (no ' ') + rtnPos := utf8.RuneCount(buf.Bytes()) + if !wordOpen { + // always required since there is a suffix + buf.WriteString(getWordOpenStr(word.Type, word.QC)) + wordOpen = true + } + buf.WriteString(string(word.Raw[wordPos:])) // write the suffix + return utilfn.StrWithPos{Str: buf.String(), Pos: rtnPos}, wordOpen +} + +// lit, grp, svar, dq, ddq, varb, sq, dsq +func Extend(word *WordType, wordPos int, extStr string, complete bool) utilfn.StrWithPos { + rtn, _ := extendInternal(word, wordPos, extStr, complete, false) + return rtn +} + +func (ec *extendContext) extend(ch rune) { + if ch == 0 { + return + } + return +} + +func isVarNameChar(ch rune) bool { + return ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || (ch >= '0' && ch <= '9') +} + +func extendVar(buf *bytes.Buffer, ch rune) { + if ch == 0 { + return + } + if !isVarNameChar(ch) { + return + } + buf.WriteRune(ch) +} + +func getSpecialEscape(ch rune) string { + if ch > unicode.MaxASCII { + return "" + } + return specialEsc[byte(ch)] +} + +func writeSpecial(buf *bytes.Buffer, ch rune, wrap bool) { + if wrap { + buf.WriteRune('$') + buf.WriteRune('\'') + } + sesc := getSpecialEscape(ch) + if sesc != "" { + buf.WriteString(sesc) + } else { + utf8Lit := getUtf8Literal(ch) + buf.WriteString(utf8Lit) + } + if wrap { + buf.WriteRune('\'') + } +} + +func extendLit(buf *bytes.Buffer, ch rune) { + if ch == 0 { + return + } + if ch > unicode.MaxASCII || !unicode.IsPrint(ch) { + writeSpecial(buf, ch, true) + return + } + var bch = byte(ch) + if noEscChars[bch] { + buf.WriteRune(ch) + return + } + buf.WriteRune('\\') + buf.WriteRune(ch) + return +} + +func extendDSQ(buf *bytes.Buffer, wordOpen *bool, ch rune) { + if ch == 0 { + return + } + if !*wordOpen { + buf.WriteRune('$') + buf.WriteRune('\'') + *wordOpen = true + } + if ch > unicode.MaxASCII || !unicode.IsPrint(ch) { + writeSpecial(buf, ch, false) + return + } + if ch == '\'' { + buf.WriteRune('\\') + buf.WriteRune(ch) + return + } + buf.WriteRune(ch) + return +} + +func extendSQ(buf *bytes.Buffer, wordOpen *bool, ch rune) { + if ch == 0 { + return + } + if ch == '\'' { + if *wordOpen { + buf.WriteRune('\'') + *wordOpen = false + } + buf.WriteRune('\\') + buf.WriteRune('\'') + return + } + if ch > unicode.MaxASCII || !unicode.IsPrint(ch) { + if *wordOpen { + buf.WriteRune('\'') + *wordOpen = false + } + writeSpecial(buf, ch, true) + return + } + if !*wordOpen { + buf.WriteRune('\'') + *wordOpen = true + } + buf.WriteRune(ch) + return +} + +func extendDQLit(buf *bytes.Buffer, wordOpen *bool, ch rune) { + if ch == 0 { + return + } + if ch > unicode.MaxASCII || !unicode.IsPrint(ch) { + if *wordOpen { + buf.WriteRune('"') + *wordOpen = false + } + writeSpecial(buf, ch, true) + return + } + if !*wordOpen { + buf.WriteRune('"') + *wordOpen = true + } + if ch == '"' || ch == '\\' || ch == '$' || ch == '`' { + buf.WriteRune('\\') + buf.WriteRune(ch) + return + } + buf.WriteRune(ch) + return +} diff --git a/wavesrv/pkg/shparse/shparse.go b/wavesrv/pkg/shparse/shparse.go new file mode 100644 index 00000000..e96fe3d9 --- /dev/null +++ b/wavesrv/pkg/shparse/shparse.go @@ -0,0 +1,693 @@ +package shparse + +import ( + "bytes" + "fmt" + + "github.com/commandlinedev/prompt-server/pkg/utilfn" +) + +// +// cmds := cmd (sep cmd)* +// sep := ';' | '&' | '&&' | '||' | '|' | '\n' +// cmd := simple-cmd | compound-command redirect-list? +// compound-command := brace-group | subshell | for-clause | case-clause | if-clause | while-clause | until-clause +// brace-group := '{' cmds '}' +// subshell := '(' cmds ')' +// simple-command := cmd-prefix cmd-word (io-redirect)* +// cmd-prefix := (io-redirect | assignment)* +// cmd-suffix := (io-redirect | word)* +// cmd-name := word +// cmd-word := word +// io-redirect := (io-number? io-file) | (io-number? io-here) +// io-file := ('<' | '<&' | '>' | '>&' | '>>' | '>|' ) filename +// io-here := ('<<' | '<<-') here_end +// here-end := word +// if-clause := 'if' compound-list 'then' compound-list else-part 'fi' +// else-part := 'elif' compound-list 'then' compound-list +// | 'elif' compount-list 'then' compound-list else-part +// | 'else' compound-list +// compound-list := linebreak term sep? +// +// +// +// A correctly-formed brace expansion must contain unquoted opening and closing braces, and at least one unquoted comma or a valid sequence expression +// Any incorrectly formed brace expansion is left unchanged. +// +// ambiguity between $((...)) and $((ls); ls) +// ambiguity between foo=([0]=hell) and foo=([abc) +// tokenization https://pubs.opengroup.org/onlinepubs/7908799/xcu/chap2.html#tag_001_003 + +// can-extend: WordTypeLit, WordTypeSimpleVar, WordTypeVarBrace, WordTypeDQ, WordTypeDDQ, WordTypeSQ, WordTypeDSQ +const ( + WordTypeRaw = "raw" + WordTypeLit = "lit" // (can-extend) + WordTypeOp = "op" // single: & ; | ( ) < > \n multi(2): && || ;; << >> <& >& <> >| (( multi(3): <<- ('((' requires special processing) + WordTypeKey = "key" // if then else elif fi do done case esac while until for in { } ! (( [[ + WordTypeGroup = "grp" // contains other words e.g. "hello"foo'bar'$x (has-subs) (can-extend) + WordTypeSimpleVar = "svar" // simplevar $ (can-extend) + + WordTypeDQ = "dq" // " (quote-context) (can-extend) (has-subs) + WordTypeDDQ = "ddq" // $" (can-extend) (has-subs) (for quotecontext, uses WordTypeDQ) + WordTypeVarBrace = "varb" // ${ (quote-context) (can-extend) (internals not parsed) + WordTypeDP = "dp" // $( (quote-context) (has-subs) + WordTypeBQ = "bq" // ` (quote-context) (has-subs) + + WordTypeSQ = "sq" // ' (can-extend) + WordTypeDSQ = "dsq" // $' (can-extend) + WordTypeDPP = "dpp" // $(( (internals not parsed) + WordTypePP = "pp" // (( (internals not parsed) + WordTypeDB = "db" // $[ (internals not parsed) +) + +const ( + CmdTypeNone = "none" // holds control structures: '(' ')' 'for' 'while' etc. + CmdTypeSimple = "simple" // holds real commands +) + +type WordType struct { + Type string + Offset int + QC QuoteContext + Raw []rune + Complete bool + Prefix []rune + Subs []*WordType +} + +type CmdType struct { + Type string + AssignmentWords []*WordType + Words []*WordType + NoneComplete bool // set to true when last-word is a "separator" +} + +type QuoteContext []string + +var wordMetaMap map[string]wordMeta + +// same order as https://www.gnu.org/software/bash/manual/html_node/Reserved-Words.html +var bashReservedWords = []string{ + "if", "then", "elif", "else", "fi", "time", + "for", "in", "until", "while", "do", "done", + "case", "esac", "coproc", "select", "function", + "{", "}", "[[", "]]", "!", +} + +// special reserved words: "for", "in", "case", "select", "function", "[[", and "]]" + +var bashNoneRW = []string{ + "if", "then", + "elif", "else", "fi", "time", + "until", "while", "do", "done", + "esac", "coproc", + "{", "}", "!", +} + +type wordMeta struct { + Type string + EmptyWord []rune + PrefixLen int + SuffixLen int + CanExtend bool + QuoteContext bool +} + +func (m wordMeta) getSuffix() string { + if m.SuffixLen == 0 { + return "" + } + return string(m.EmptyWord[len(m.EmptyWord)-m.SuffixLen:]) +} + +func (m wordMeta) getPrefix() string { + if m.PrefixLen == 0 { + return "" + } + return string(m.EmptyWord[:m.PrefixLen]) +} + +func makeWordMeta(wtype string, emptyWord string, prefixLen int, suffixLen int, canExtend bool, quoteContext bool) { + if len(emptyWord) != prefixLen+suffixLen { + panic(fmt.Sprintf("invalid empty word %s %d %d", emptyWord, prefixLen, suffixLen)) + } + wordMetaMap[wtype] = wordMeta{wtype, []rune(emptyWord), prefixLen, suffixLen, canExtend, quoteContext} +} + +func init() { + wordMetaMap = make(map[string]wordMeta) + makeWordMeta(WordTypeRaw, "", 0, 0, false, false) + makeWordMeta(WordTypeLit, "", 0, 0, true, false) + makeWordMeta(WordTypeOp, "", 0, 0, false, false) + makeWordMeta(WordTypeKey, "", 0, 0, false, false) + makeWordMeta(WordTypeGroup, "", 0, 0, false, false) + makeWordMeta(WordTypeSimpleVar, "$", 1, 0, true, false) + makeWordMeta(WordTypeVarBrace, "${}", 2, 1, true, true) + makeWordMeta(WordTypeDQ, `""`, 1, 1, true, true) + makeWordMeta(WordTypeDDQ, `$""`, 2, 1, true, true) + makeWordMeta(WordTypeDP, "$()", 2, 1, false, false) + makeWordMeta(WordTypeBQ, "``", 1, 1, false, false) + makeWordMeta(WordTypeSQ, "''", 1, 1, true, false) + makeWordMeta(WordTypeDSQ, "$''", 2, 1, true, false) + makeWordMeta(WordTypeDPP, "$(())", 3, 2, false, false) + makeWordMeta(WordTypePP, "(())", 2, 2, false, false) + makeWordMeta(WordTypeDB, "$[]", 2, 1, false, false) +} + +func MakeEmptyWord(wtype string, qc QuoteContext, offset int, complete bool) *WordType { + meta := wordMetaMap[wtype] + if meta.Type == "" { + meta = wordMetaMap[WordTypeRaw] + } + rtn := &WordType{Type: meta.Type, QC: qc, Offset: offset, Complete: complete} + if len(meta.EmptyWord) > 0 { + if complete { + rtn.Raw = append([]rune(nil), meta.EmptyWord...) + } else { + rtn.Raw = append([]rune(nil), []rune(meta.getPrefix())...) + } + } + return rtn +} + +func (qc QuoteContext) push(q string) QuoteContext { + rtn := make([]string, 0, len(qc)+1) + rtn = append(rtn, qc...) + rtn = append(rtn, q) + return rtn +} + +func (qc QuoteContext) cur() string { + if len(qc) == 0 { + return "" + } + return qc[len(qc)-1] +} + +func (qc QuoteContext) clone() QuoteContext { + if len(qc) == 0 { + return nil + } + return append([]string(nil), qc...) +} + +func makeRepeatStr(ch byte, slen int) string { + if slen == 0 { + return "" + } + rtn := make([]byte, slen) + for i := 0; i < slen; i++ { + rtn[i] = ch + } + return string(rtn) +} + +func (w *WordType) isBlank() bool { + return w.Type == WordTypeLit && len(w.Raw) == 0 +} + +func (w *WordType) contentEndPos() int { + if !w.Complete { + return len(w.Raw) + } + wmeta := wordMetaMap[w.Type] + return len(w.Raw) - wmeta.SuffixLen +} + +func (w *WordType) contentStartPos() int { + wmeta := wordMetaMap[w.Type] + return wmeta.PrefixLen +} + +func (w *WordType) canHaveSubs() bool { + switch w.Type { + case WordTypeGroup, WordTypeDQ, WordTypeDDQ, WordTypeDP, WordTypeBQ: + return true + + default: + return false + } +} + +func (w *WordType) uncompletable() bool { + switch w.Type { + case WordTypeRaw, WordTypeOp, WordTypeKey, WordTypeDPP, WordTypePP, WordTypeDB, WordTypeBQ, WordTypeDP: + return true + + default: + return false + } +} + +func (w *WordType) stringWithPos(pos int) string { + notCompleteFlag := " " + if !w.Complete { + notCompleteFlag = "*" + } + str := string(w.Raw) + if pos != -1 { + str = utilfn.StrWithPos{Str: str, Pos: pos}.String() + } + return fmt.Sprintf("%-4s[%3d]%s %s%q", w.Type, w.Offset, notCompleteFlag, makeRepeatStr('_', len(w.Prefix)), str) +} + +func (w *WordType) String() string { + notCompleteFlag := " " + if !w.Complete { + notCompleteFlag = "*" + } + return fmt.Sprintf("%-4s[%3d]%s %s%q", w.Type, w.Offset, notCompleteFlag, makeRepeatStr('_', len(w.Prefix)), string(w.Raw)) +} + +// offset = -1 for don't show +func dumpWords(words []*WordType, indentStr string, offset int) { + wrotePos := false + for _, word := range words { + posInWord := false + if !wrotePos && offset != -1 && offset <= word.Offset { + fmt.Printf("%s* [%3d] [*]\n", indentStr, offset) + wrotePos = true + } + if !wrotePos && offset != -1 && offset < word.Offset+len(word.Raw) { + fmt.Printf("%s%s\n", indentStr, word.stringWithPos(offset-word.Offset)) + wrotePos = true + posInWord = true + } else { + fmt.Printf("%s%s\n", indentStr, word.String()) + } + if len(word.Subs) > 0 { + if posInWord { + wmeta := wordMetaMap[word.Type] + dumpWords(word.Subs, indentStr+" ", offset-word.Offset-wmeta.PrefixLen) + } else { + dumpWords(word.Subs, indentStr+" ", -1) + } + } + } +} + +func dumpCommands(cmds []*CmdType, indentStr string, pos int) { + for _, cmd := range cmds { + fmt.Printf("%sCMD: %s [%d] pos:%d\n", indentStr, cmd.Type, len(cmd.Words), pos) + dumpWords(cmd.AssignmentWords, indentStr+" *", pos) + dumpWords(cmd.Words, indentStr+" ", pos) + } +} + +func wordsToStr(words []*WordType) string { + var buf bytes.Buffer + for _, word := range words { + if len(word.Prefix) > 0 { + buf.WriteString(string(word.Prefix)) + } + buf.WriteString(string(word.Raw)) + } + return buf.String() +} + +// recognizes reserved words in first position +func convertToAnyReservedWord(w *WordType) bool { + if w == nil || w.Type != WordTypeLit { + return false + } + rawVal := string(w.Raw) + for _, rw := range bashReservedWords { + if rawVal == rw { + w.Type = WordTypeKey + return true + } + } + return false +} + +// recognizes the specific reserved-word given only ('in' and 'do' in 'for', 'case', and 'select' commands) +func convertToReservedWord(w *WordType, reservedWord string) { + if w == nil || w.Type != WordTypeLit { + return + } + if string(w.Raw) == reservedWord { + w.Type = WordTypeKey + } +} + +func isNoneReservedWord(w *WordType) bool { + if w.Type != WordTypeKey { + return false + } + rawVal := string(w.Raw) + for _, rw := range bashNoneRW { + if rawVal == rw { + return true + } + } + return false +} + +type parseCmdState struct { + Input []*WordType + InputPos int + + Rtn []*CmdType + Cur *CmdType +} + +func (state *parseCmdState) isEof() bool { + return state.InputPos >= len(state.Input) +} + +func (state *parseCmdState) curWord() *WordType { + if state.isEof() { + return nil + } + return state.Input[state.InputPos] +} + +func (state *parseCmdState) lastCmd() *CmdType { + if len(state.Rtn) == 0 { + return nil + } + return state.Rtn[len(state.Rtn)-1] +} + +func (state *parseCmdState) makeNoneCmd(sep bool) { + if state.Cur == nil || state.Cur.Type != CmdTypeNone { + state.Cur = &CmdType{Type: CmdTypeNone} + state.Rtn = append(state.Rtn, state.Cur) + } + state.Cur.Words = append(state.Cur.Words, state.curWord()) + if sep { + state.Cur.NoneComplete = true + state.Cur = nil + } + state.InputPos++ +} + +func (state *parseCmdState) handleKeyword(word *WordType) bool { + if word.Type != WordTypeKey { + return false + } + if isNoneReservedWord(word) { + state.makeNoneCmd(true) + return true + } + rw := string(word.Raw) + if rw == "[[" { + // just ignore everything between [[ and ]] + for !state.isEof() { + curWord := state.curWord() + if curWord.Type == WordTypeLit && string(curWord.Raw) == "]]" { + convertToReservedWord(curWord, "]]") + state.makeNoneCmd(false) + break + } + state.makeNoneCmd(false) + } + return true + } + if rw == "case" { + // ignore everything between "case" and "esac" + for !state.isEof() { + curWord := state.curWord() + if curWord.Type == WordTypeKey && string(curWord.Raw) == "esac" { + state.makeNoneCmd(false) + break + } + state.makeNoneCmd(false) + } + return true + } + if rw == "for" || rw == "select" { + // ignore until a "do" + for !state.isEof() { + curWord := state.curWord() + if curWord.Type == WordTypeKey && string(curWord.Raw) == "do" { + state.makeNoneCmd(true) + break + } + state.makeNoneCmd(false) + } + return true + } + if rw == "in" { + // the "for" and "case" clauses should skip "in". so encountering an "in" here is a syntax error. + // just treat it as a none and allow a new command after. + state.makeNoneCmd(false) + return true + } + if rw == "function" { + // ignore until '{' + for !state.isEof() { + curWord := state.curWord() + if curWord.Type == WordTypeKey && string(curWord.Raw) == "{" { + state.makeNoneCmd(true) + break + } + state.makeNoneCmd(false) + } + return true + } + state.makeNoneCmd(true) + return true +} + +func isCmdSeparatorOp(word *WordType) bool { + if word.Type != WordTypeOp { + return false + } + opVal := string(word.Raw) + return opVal == ";" || opVal == "\n" || opVal == "&" || opVal == "|" || opVal == "|&" || opVal == "&&" || opVal == "||" || opVal == "(" || opVal == ")" +} + +func (state *parseCmdState) handleOp(word *WordType) bool { + opVal := string(word.Raw) + // sequential separators + if opVal == ";" || opVal == "\n" { + state.makeNoneCmd(true) + return true + } + // separator + if opVal == "&" { + state.makeNoneCmd(true) + return true + } + // pipelines + if opVal == "|" || opVal == "|&" { + state.makeNoneCmd(true) + return true + } + // lists + if opVal == "&&" || opVal == "||" { + state.makeNoneCmd(true) + return true + } + // subshell + if opVal == "(" || opVal == ")" { + state.makeNoneCmd(true) + return true + } + return false +} + +func wordSliceBoundedIdx(words []*WordType, idx int) *WordType { + if idx >= len(words) { + return nil + } + return words[idx] +} + +// note that a newline "op" can appear in the third position of "for" or "case". the "in" keyword is still converted because of wordNum == 0 +func identifyReservedWords(words []*WordType) { + wordNum := 0 + lastReserved := false + for idx, word := range words { + if wordNum == 0 || lastReserved { + convertToAnyReservedWord(word) + } + if word.Type == WordTypeKey { + rwVal := string(word.Raw) + switch rwVal { + case "for": + lastReserved = false + third := wordSliceBoundedIdx(words, idx+2) + convertToReservedWord(third, "in") + convertToReservedWord(third, "do") + + case "case": + lastReserved = false + third := wordSliceBoundedIdx(words, idx+2) + convertToReservedWord(third, "in") + + case "in": + lastReserved = false + + default: + lastReserved = true + } + continue + } + lastReserved = false + if isCmdSeparatorOp(word) { + wordNum = 0 + continue + } + wordNum++ + } +} + +func ResetWordOffsets(words []*WordType, startIdx int) { + pos := startIdx + for _, word := range words { + pos += len(word.Prefix) + word.Offset = pos + if len(word.Subs) > 0 { + ResetWordOffsets(word.Subs, 0) + } + pos += len(word.Raw) + } +} + +func CommandsToWords(cmds []*CmdType) []*WordType { + var rtn []*WordType + for _, cmd := range cmds { + rtn = append(rtn, cmd.Words...) + } + return rtn +} + +func (c *CmdType) stripPrefix() []rune { + if len(c.AssignmentWords) > 0 { + w := c.AssignmentWords[0] + prefix := w.Prefix + if len(prefix) == 0 { + return nil + } + newWord := *w + newWord.Prefix = nil + c.AssignmentWords[0] = &newWord + return prefix + } + if len(c.Words) > 0 { + w := c.Words[0] + prefix := w.Prefix + if len(prefix) == 0 { + return nil + } + newWord := *w + newWord.Prefix = nil + c.Words[0] = &newWord + return prefix + } + return nil +} + +func (c *CmdType) isEmpty() bool { + return len(c.AssignmentWords) == 0 && len(c.Words) == 0 +} + +func (c *CmdType) lastWord() *WordType { + if len(c.Words) > 0 { + return c.Words[len(c.Words)-1] + } + if len(c.AssignmentWords) > 0 { + return c.AssignmentWords[len(c.AssignmentWords)-1] + } + return nil +} + +func (c *CmdType) firstWord() *WordType { + if len(c.AssignmentWords) > 0 { + return c.AssignmentWords[0] + } + if len(c.Words) > 0 { + return c.Words[0] + } + return nil +} + +func (c *CmdType) offset() int { + firstWord := c.firstWord() + if firstWord == nil { + return 0 + } + return firstWord.Offset +} + +func (c *CmdType) endOffset() int { + lastWord := c.lastWord() + if lastWord == nil { + return 0 + } + return lastWord.Offset + len(lastWord.Raw) +} + +func indexInRunes(arr []rune, ch rune) int { + for idx, r := range arr { + if r == ch { + return idx + } + } + return -1 +} + +func isAssignmentWord(w *WordType) bool { + if w.Type == WordTypeLit || w.Type == WordTypeGroup { + eqIdx := indexInRunes(w.Raw, '=') + if eqIdx == -1 { + return false + } + prefix := w.Raw[0:eqIdx] + return isSimpleVarName(prefix) + } + return false +} + +// simple commands steal whitespace from subsequent commands +func cmdWhitespaceFixup(cmds []*CmdType) { + for idx := 0; idx < len(cmds)-1; idx++ { + cmd := cmds[idx] + if cmd.Type != CmdTypeSimple || cmd.isEmpty() { + continue + } + nextCmd := cmds[idx+1] + nextPrefix := nextCmd.stripPrefix() + if len(nextPrefix) > 0 { + blankWord := &WordType{Type: WordTypeLit, QC: cmd.lastWord().QC, Offset: cmd.endOffset() + len(nextPrefix), Prefix: nextPrefix, Complete: true} + cmd.Words = append(cmd.Words, blankWord) + } + } +} + +func ParseCommands(words []*WordType) []*CmdType { + identifyReservedWords(words) + state := parseCmdState{Input: words} + for { + if state.isEof() { + break + } + word := state.curWord() + if word.Type == WordTypeKey { + done := state.handleKeyword(word) + if done { + continue + } + } + if word.Type == WordTypeOp { + done := state.handleOp(word) + if done { + continue + } + } + if state.Cur == nil || state.Cur.Type != CmdTypeSimple { + state.Cur = &CmdType{Type: CmdTypeSimple} + state.Rtn = append(state.Rtn, state.Cur) + } + if len(state.Cur.Words) == 0 && isAssignmentWord(word) { + state.Cur.AssignmentWords = append(state.Cur.AssignmentWords, word) + } else { + state.Cur.Words = append(state.Cur.Words, word) + } + state.InputPos++ + } + cmdWhitespaceFixup(state.Rtn) + return state.Rtn +} diff --git a/wavesrv/pkg/shparse/shparse_test.go b/wavesrv/pkg/shparse/shparse_test.go new file mode 100644 index 00000000..16c8a83a --- /dev/null +++ b/wavesrv/pkg/shparse/shparse_test.go @@ -0,0 +1,219 @@ +package shparse + +import ( + "fmt" + "testing" + + "github.com/commandlinedev/prompt-server/pkg/utilfn" +) + +// $(ls f[*]); ./x +// ls f => raw["ls f"] -> lit["ls f"] -> lit["ls"] lit["f"] +// w; ls foo; => raw["w; ls foo;"] +// ls&"ls" => raw["ls&ls"] => lit["ls&"] dq["ls"] => lit["ls"] key["&"] dq["ls"] +// ls $x; echo `ls f => raw["ls $x; echo `ls f"] +// > echo $foo{x,y} + +func testParse(t *testing.T, s string) { + words := Tokenize(s) + + fmt.Printf("parse <<\n%s\n>>\n", s) + dumpWords(words, " ", 8) + outStr := wordsToStr(words) + if outStr != s { + t.Errorf("tokenization output does not match input: %q => %q", s, outStr) + } + fmt.Printf("------\n\n") +} + +func Test1(t *testing.T) { + testParse(t, "ls") + testParse(t, "ls 'foo'") + testParse(t, `ls "hello" $'\''`) + testParse(t, `ls "foo`) + testParse(t, `echo $11 $xyz $ `) + testParse(t, `echo $(ls ${x:"hello"} foo`) + testParse(t, `ls ${x:"hello"} $[2+2] $((5 * 10)) $(ls; ls&)`) + testParse(t, `ls;ls&./foo > out 2> "out2"`) + testParse(t, `(( x = 5)); ls& cd ~/work/"hello again"`) + testParse(t, `echo "hello"abc$(ls)$x${y:foo}`) + testParse(t, `echo $(ls; ./x "foo")`) + testParse(t, `echo $(ls; (cd foo; ls); (cd bar; ls))xyz`) + testParse(t, `echo "$x ${y:-foo}"`) + testParse(t, `command="$(echo "$input" | sed -e "s/^[ \t]*\([^ \t]*\)[ \t]*.*$/\1/g")"`) + testParse(t, `echo $(ls $)`) + testParse(t, `echo ${x:-hello\}"}"} 2nd`) + testParse(t, `echo "$(ls "foo") more $x"`) + testParse(t, "echo `ls $x \"hello $x\" \\`ls\\`; ./foo`") + testParse(t, `echo $"hello $x $(ls)"`) + testParse(t, "echo 'hello'\nls\n") + testParse(t, "echo 'hello'abc$'\a'") +} + +func lastWord(words []*WordType) *WordType { + if len(words) == 0 { + return nil + } + return words[len(words)-1] +} + +func testExtend(t *testing.T, startStr string, extendStr string, complete bool, expStr string) { + startSP := utilfn.ParseToSP(startStr) + words := Tokenize(startSP.Str) + word := findCompletionWordAtPos(words, startSP.Pos, true) + if word == nil { + word = MakeEmptyWord(WordTypeLit, nil, startSP.Pos, true) + } + outSP := Extend(word, startSP.Pos-word.Offset, extendStr, complete) + expSP := utilfn.ParseToSP(expStr) + fmt.Printf("extend: [%s] + %q => [%s]\n", startStr, extendStr, outSP) + if outSP != expSP { + t.Errorf("extension does not match: [%s] + %q => [%s] expected [%s]\n", startStr, extendStr, outSP, expSP) + } +} + +func Test2(t *testing.T) { + testExtend(t, `he[*]`, "llo", false, "hello[*]") + testExtend(t, `he[*]`, "llo", true, "hello [*]") + testExtend(t, `'mi[*]e`, "k", false, "'mik[*]e") + testExtend(t, `'mi[*]e`, "k", true, "'mik[*]e") + testExtend(t, `'mi[*]'`, "ke", true, "'mike' [*]") + testExtend(t, `'mi'[*]`, "ke", true, "'mike' [*]") + testExtend(t, `'mi[*]'`, "ke", false, "'mike[*]'") + testExtend(t, `'mi'[*]`, "ke", false, "'mike[*]'") + testExtend(t, `$f[*]`, "oo", false, "$foo[*]") + testExtend(t, `${f}[*]`, "oo", false, "${foo[*]}") + testExtend(t, `${f[*]}`, "oo", true, "${foo} [*]") + testExtend(t, `[*]`, "more stuff", false, `more\ stuff[*]`) + testExtend(t, `[*]`, "hello\amike", false, `hello$'\a'mike[*]`) + testExtend(t, `$'he[*]'`, "\x01\x02\x0a", true, `$'he\x01\x02\n' [*]`) + testExtend(t, `${x}\ [*]ll$y`, "e", false, `${x}\ e[*]ll$y`) + testExtend(t, `"he[*]"`, "$$o", true, `"he\$\$o" [*]`) + testExtend(t, `"h[*]llo"`, "e", false, `"he[*]llo"`) + testExtend(t, `"h[*]llo"`, "e", true, `"he[*]llo"`) + testExtend(t, `"[*]${h}llo"`, "e\x01", true, `"e"$'\x01'[*]"${h}llo"`) + testExtend(t, `"${h}llo[*]"`, "e\x01", true, `"${h}lloe"$'\x01' [*]`) + testExtend(t, `"${h}llo[*]"`, "e\x01", false, `"${h}lloe"$'\x01'[*]`) + testExtend(t, `"${h}ll[*]o"`, "e\x01", false, `"${h}lle"$'\x01'[*]"o"`) + testExtend(t, `"ab[*]c${x}def"`, "\x01", false, `"ab"$'\x01'[*]"c${x}def"`) + testExtend(t, `'ab[*]ef'`, "\x01", false, `'ab'$'\x01'[*]'ef'`) + + // testExtend(t, `'he'`, "llo", `'hello'`) + // testExtend(t, `'he'`, "'", `'he'\'''`) + // testExtend(t, `'he'`, "'\x01", `'he'\'$'\x01'''`) + // testExtend(t, `he`, "llo", `hello`) + // testExtend(t, `he`, "l*l'\x01\x07o", `hel\*l\'$'\x01'$'\a'o`) + // testExtend(t, `$x`, "fo|o", `$xfoo`) + // testExtend(t, `${x`, "fo|o", `${xfoo`) + // testExtend(t, `$'f`, "oo", `$'foo`) + // testExtend(t, `$'f`, "'\x01\x07o", `$'f\'\x01\ao`) + // testExtend(t, `"f"`, "oo", `"foo"`) + // testExtend(t, `"mi"`, "ke's \"hello\"", `"mike's \"hello\""`) + // testExtend(t, `"t"`, "t\x01\x07", `"tt"$'\x01'$'\a'""`) +} + +func testParseCommands(t *testing.T, str string) { + fmt.Printf("parse: %q\n", str) + words := Tokenize(str) + cmds := ParseCommands(words) + dumpCommands(cmds, " ", -1) + fmt.Printf("\n") +} + +func TestCmd(t *testing.T) { + testParseCommands(t, "ls foo") + testParseCommands(t, "function foo () { echo hello; }") + testParseCommands(t, "ls foo && ls bar; ./run $x hello | xargs foo; ") + testParseCommands(t, "if [[ 2 > 1 ]]; then echo hello\nelse echo world; echo next; done") + testParseCommands(t, "case lots of stuff; i don\\'t know how to parse; esac; ls foo") + testParseCommands(t, "(ls & ./x \n \n); for x in $vars 3; do { echo $x; ls foo ; } done") + testParseCommands(t, `ls f"oo" "${x:"hello$y"}"`) + testParseCommands(t, `x="foo $y" z=10 ls`) +} + +func testCompPos(t *testing.T, cmdStr string, compType string, hasCommand bool, cmdWordPos int, hasWord bool, superOffset int) { + cmdSP := utilfn.ParseToSP(cmdStr) + words := Tokenize(cmdSP.Str) + cmds := ParseCommands(words) + cpos := FindCompletionPos(cmds, cmdSP.Pos) + fmt.Printf("testCompPos [%d] %q => [%s] %v\n", cmdSP.Pos, cmdStr, cpos.CompType, cpos) + if cpos.CompType != compType { + t.Errorf("testCompPos %q => invalid comp-type %q, expected %q", cmdStr, cpos.CompType, compType) + } + if cpos.CompWord != nil { + fmt.Printf(" found-word: %d %s\n", cpos.CompWordOffset, cpos.CompWord.stringWithPos(cpos.CompWordOffset)) + } + if cpos.Cmd != nil { + fmt.Printf(" found-cmd: ") + dumpCommands([]*CmdType{cpos.Cmd}, " ", cpos.RawPos) + } + dumpCommands(cmds, " ", cmdSP.Pos) + fmt.Printf("\n") + if cpos.RawPos+cpos.SuperOffset != cmdSP.Pos { + t.Errorf("testCompPos %q => bad rawpos:%d superoffset:%d expected:%d", cmdStr, cpos.RawPos, cpos.SuperOffset, cmdSP.Pos) + } + if (cpos.Cmd != nil) != hasCommand { + t.Errorf("testCompPos %q => bad has-command exp:%v", cmdStr, hasCommand) + } + if (cpos.CompWord != nil) != hasWord { + t.Errorf("testCompPos %q => bad has-word exp:%v", cmdStr, hasWord) + } + if cpos.CmdWordPos != cmdWordPos { + t.Errorf("testCompPos %q => bad cmd-word-pos got:%d exp:%d", cmdStr, cpos.CmdWordPos, cmdWordPos) + } + if cpos.SuperOffset != superOffset { + t.Errorf("testCompPos %q => bad super-offset got:%d exp:%d", cmdStr, cpos.SuperOffset, superOffset) + } +} + +func TestCompPos(t *testing.T) { + testCompPos(t, "ls [*]foo", CompTypeArg, true, 1, false, 0) + testCompPos(t, "ls foo [*];", CompTypeArg, true, 2, false, 0) + testCompPos(t, "ls foo ;[*]", CompTypeCommand, false, 0, false, 0) + testCompPos(t, "ls foo >[*]> ./bar", CompTypeInvalid, true, 2, true, 0) + testCompPos(t, "l[*]s", CompTypeCommand, true, 0, true, 0) + testCompPos(t, "ls[*]", CompTypeCommand, true, 0, true, 0) + testCompPos(t, "x=10 { (ls ./f[*] more); ls }", CompTypeArg, true, 1, true, 0) + testCompPos(t, "for x in 1[*] 2 3; do ", CompTypeBasic, false, 0, true, 0) + testCompPos(t, "for[*] x in 1 2 3;", CompTypeInvalid, false, 0, true, 0) + testCompPos(t, `ls "abc $(ls -l t[*])" && foo`, CompTypeArg, true, 2, true, 10) + testCompPos(t, "ls ${abc:$(ls -l [*])}", CompTypeVar, false, 0, true, 0) // we don't sub-parse inside of ${} (so this returns "var" right now) + testCompPos(t, `ls abc"$(ls $"echo $(ls ./[*]x) foo)" `, CompTypeArg, true, 1, true, 21) + testCompPos(t, `ls "abc$d[*]"`, CompTypeVar, false, 0, true, 4) + testCompPos(t, `ls "abc$d$'a[*]`, CompTypeArg, true, 1, true, 0) + testCompPos(t, `ls $[*]'foo`, CompTypeArg, true, 1, true, 0) + testCompPos(t, `echo $TE[*]`, CompTypeVar, false, 0, true, 0) +} + +func testExpand(t *testing.T, str string, pos int, expStr string, expInfo *ExpandInfo) { + ectx := ExpandContext{HomeDir: "/Users/mike"} + words := Tokenize(str) + if len(words) == 0 { + t.Errorf("could not tokenize any words from %q", str) + return + } + word := words[0] + output, info := SimpleExpandPrefix(ectx, word, pos) + if output != expStr { + t.Errorf("error expanding %q, output:%q exp:%q", str, output, expStr) + } else { + fmt.Printf("expand: %q (%d) => %q\n", str, pos, output) + } + if expInfo != nil { + if info != *expInfo { + t.Errorf("error expanding %q, info:%v exp:%v", str, info, expInfo) + } + } +} + +func TestExpand(t *testing.T) { + testExpand(t, "hello", 3, "hel", nil) + testExpand(t, "he\\$xabc", 6, "he$xa", nil) + testExpand(t, "he${x}abc", 6, "he${x}", nil) + testExpand(t, "'hello\"mike'", 8, "hello\"m", nil) + testExpand(t, `$'abc\x01def`, 10, "abc\x01d", nil) + testExpand(t, `$((2 + 2))`, 6, "$((2 +", &ExpandInfo{HasSpecial: true}) + testExpand(t, `abc"def"`, 6, "abcde", nil) + testExpand(t, `"abc$x$'"'""`, 12, "abc$x\"", nil) + testExpand(t, `'he'\''s'`, 9, "he's", nil) +} diff --git a/wavesrv/pkg/shparse/tokenize.go b/wavesrv/pkg/shparse/tokenize.go new file mode 100644 index 00000000..825e7387 --- /dev/null +++ b/wavesrv/pkg/shparse/tokenize.go @@ -0,0 +1,601 @@ +package shparse + +import ( + "fmt" + "unicode" +) + +// from bash source +// +// shell_meta_chars "()<>;&|" +// + +type tokenizeOutputState struct { + Rtn []*WordType + CurWord *WordType + SavedPrefix []rune +} + +func copyRunes(rarr []rune) []rune { + if len(rarr) == 0 { + return nil + } + return append([]rune(nil), rarr...) +} + +// does not set CurWord +func (state *tokenizeOutputState) appendStandaloneWord(word *WordType) { + state.delimitCurWord() + if len(state.SavedPrefix) > 0 { + word.Prefix = state.SavedPrefix + state.SavedPrefix = nil + } + state.Rtn = append(state.Rtn, word) +} + +func (state *tokenizeOutputState) appendWord(word *WordType) { + if len(state.SavedPrefix) > 0 { + word.Prefix = state.SavedPrefix + state.SavedPrefix = nil + } + if state.CurWord == nil { + state.CurWord = word + return + } + state.ensureGroupWord() + word.Offset = word.Offset - state.CurWord.Offset + state.CurWord.Subs = append(state.CurWord.Subs, word) + state.CurWord.Raw = append(state.CurWord.Raw, word.Raw...) +} + +func (state *tokenizeOutputState) ensureGroupWord() { + if state.CurWord == nil { + panic("invalid state, cannot make group word when CurWord is nil") + } + if state.CurWord.Type == WordTypeGroup { + return + } + // moves the prefix from CurWord to the new group word, resets offsets + groupWord := &WordType{ + Type: WordTypeGroup, + Offset: state.CurWord.Offset, + QC: state.CurWord.QC, + Raw: copyRunes(state.CurWord.Raw), + Complete: true, + Prefix: state.CurWord.Prefix, + } + state.CurWord.Prefix = nil + state.CurWord.Offset = 0 + groupWord.Subs = []*WordType{state.CurWord} + state.CurWord = groupWord +} + +func ungroupWord(groupWord *WordType) []*WordType { + if groupWord.Type != WordTypeGroup { + return []*WordType{groupWord} + } + rtn := groupWord.Subs + if len(groupWord.Prefix) > 0 && len(rtn) > 0 { + newPrefix := append([]rune{}, groupWord.Prefix...) + newPrefix = append(newPrefix, rtn[0].Prefix...) + rtn[0].Prefix = newPrefix + } + for _, word := range rtn { + word.Offset = word.Offset + groupWord.Offset + } + return rtn +} + +func (state *tokenizeOutputState) ensureLitCurWord(pc *parseContext) { + if state.CurWord == nil { + state.CurWord = pc.makeWord(WordTypeLit, 0, true) + state.CurWord.Prefix = state.SavedPrefix + state.SavedPrefix = nil + return + } + if state.CurWord.Type == WordTypeLit { + return + } + state.ensureGroupWord() + lastWord := state.CurWord.Subs[len(state.CurWord.Subs)-1] + if lastWord.Type != WordTypeLit { + if len(state.SavedPrefix) > 0 { + panic("invalid state, there can be no saved prefix") + } + litWord := pc.makeWord(WordTypeLit, 0, true) + litWord.Offset = litWord.Offset - state.CurWord.Offset + state.CurWord.Subs = append(state.CurWord.Subs, litWord) + } +} + +func (state *tokenizeOutputState) delimitCurWord() { + if state.CurWord != nil { + state.Rtn = append(state.Rtn, state.CurWord) + state.CurWord = nil + } +} + +func (state *tokenizeOutputState) delimitWithSpace(spaceCh rune) { + state.delimitCurWord() + state.SavedPrefix = append(state.SavedPrefix, spaceCh) +} + +func (state *tokenizeOutputState) appendLiteral(pc *parseContext, ch rune) { + state.ensureLitCurWord(pc) + if state.CurWord.Type == WordTypeLit { + state.CurWord.Raw = append(state.CurWord.Raw, ch) + } else if state.CurWord.Type == WordTypeGroup { + lastWord := state.CurWord.Subs[len(state.CurWord.Subs)-1] + if lastWord.Type != WordTypeLit { + panic(fmt.Sprintf("invalid curword type (group) %q", state.CurWord.Type)) + } + lastWord.Raw = append(lastWord.Raw, ch) + state.CurWord.Raw = append(state.CurWord.Raw, ch) + } else { + panic(fmt.Sprintf("invalid curword type %q", state.CurWord.Type)) + } +} + +func (state *tokenizeOutputState) finish(pc *parseContext) { + state.delimitCurWord() + if len(state.SavedPrefix) > 0 { + state.ensureLitCurWord(pc) + state.delimitCurWord() + } +} + +func (c *parseContext) tokenizeVarBrace() ([]*WordType, bool) { + state := &tokenizeOutputState{} + eofExit := false + for { + ch := c.cur() + if ch == 0 { + eofExit = true + break + } + if ch == '}' { + c.Pos++ + break + } + var quoteWord *WordType + if ch == '\'' { + quoteWord = c.parseStrSQ() + } + if quoteWord == nil && ch == '"' { + quoteWord = c.parseStrDQ() + } + isNextBrace := c.at(1) == '}' + if quoteWord == nil && ch == '$' && !isNextBrace { + quoteWord = c.parseStrANSI() + if quoteWord == nil { + quoteWord = c.parseStrDDQ() + } + if quoteWord == nil { + quoteWord = c.parseExpansion() + } + } + if quoteWord != nil { + state.appendWord(quoteWord) + continue + } + if ch == '\\' && c.at(1) != 0 { + state.appendLiteral(c, ch) + state.appendLiteral(c, c.at(1)) + c.Pos += 2 + continue + } + state.appendLiteral(c, ch) + c.Pos++ + } + return state.Rtn, eofExit +} + +func (c *parseContext) tokenizeDQ() ([]*WordType, bool) { + state := &tokenizeOutputState{} + eofExit := false + for { + ch := c.cur() + if ch == 0 { + eofExit = true + break + } + if ch == '"' { + c.Pos++ + break + } + if ch == '$' && c.at(1) != 0 { + quoteWord := c.parseStrANSI() + if quoteWord == nil { + quoteWord = c.parseStrDDQ() + } + if quoteWord == nil { + quoteWord = c.parseExpansion() + } + if quoteWord != nil { + state.appendWord(quoteWord) + continue + } + } + if ch == '\\' && c.at(1) != 0 { + state.appendLiteral(c, ch) + state.appendLiteral(c, c.at(1)) + c.Pos += 2 + continue + } + state.appendLiteral(c, ch) + c.Pos++ + } + state.finish(c) + if len(state.Rtn) == 0 { + return nil, eofExit + } + if len(state.Rtn) == 1 && state.Rtn[0].Type == WordTypeGroup { + return ungroupWord(state.Rtn[0]), eofExit + } + return state.Rtn, eofExit +} + +// returns (words, eofexit) +// backticks (WordTypeBQ) handle backslash in a special way, but that seems to mainly effect execution (not completion) +// de_backslash => removes initial backslash in \`, \\, and \$ before execution +func (c *parseContext) tokenizeRaw() ([]*WordType, bool) { + state := &tokenizeOutputState{} + isExpSubShell := c.QC.cur() == WordTypeDP + isInBQ := c.QC.cur() == WordTypeBQ + parenLevel := 0 + eofExit := false + for { + ch := c.cur() + if ch == 0 { + eofExit = true + break + } + if isExpSubShell && ch == ')' && parenLevel == 0 { + c.Pos++ + break + } + if isInBQ && ch == '`' { + c.Pos++ + break + } + // fmt.Printf("ch %d %q\n", c.Pos, string([]rune{ch})) + foundOp, newOffset := c.parseOp(0) + if foundOp { + opVal := string(c.Input[c.Pos : c.Pos+newOffset]) + if opVal == "(" { + arithWord := c.parseArith(true) + if arithWord != nil { + state.appendStandaloneWord(arithWord) + continue + } else { + parenLevel++ + } + } + if opVal == ")" { + parenLevel-- + } + opWord := c.makeWord(WordTypeOp, newOffset, true) + state.appendStandaloneWord(opWord) + continue + } + var quoteWord *WordType + if ch == '\'' { + quoteWord = c.parseStrSQ() + } + if quoteWord == nil && ch == '"' { + quoteWord = c.parseStrDQ() + } + if quoteWord == nil && ch == '`' { + quoteWord = c.parseStrBQ() + } + isNextParen := isExpSubShell && c.at(1) == ')' + if quoteWord == nil && ch == '$' && !isNextParen { + quoteWord = c.parseStrANSI() + if quoteWord == nil { + quoteWord = c.parseStrDDQ() + } + if quoteWord == nil { + quoteWord = c.parseExpansion() + } + } + if quoteWord != nil { + state.appendWord(quoteWord) + continue + } + if ch == '\\' && c.at(1) != 0 { + state.appendLiteral(c, ch) + state.appendLiteral(c, c.at(1)) + c.Pos += 2 + continue + } + if ch == '\n' { + newlineWord := c.makeWord(WordTypeOp, 1, true) + state.appendStandaloneWord(newlineWord) + continue + } + if unicode.IsSpace(ch) { + state.delimitWithSpace(ch) + c.Pos++ + continue + } + state.appendLiteral(c, ch) + c.Pos++ + } + state.finish(c) + return state.Rtn, eofExit +} + +type parseContext struct { + Input []rune + Pos int + QC QuoteContext +} + +func (c *parseContext) clone(pos int, newQuote string) *parseContext { + rtn := parseContext{Input: c.Input[pos:], QC: c.QC} + if newQuote != "" { + rtn.QC = rtn.QC.push(newQuote) + } + return &rtn +} + +func (c *parseContext) at(offset int) rune { + pos := c.Pos + offset + if pos < 0 || pos >= len(c.Input) { + return 0 + } + return c.Input[pos] +} + +func (c *parseContext) eof() bool { + return c.Pos >= len(c.Input) +} + +func (c *parseContext) cur() rune { + return c.at(0) +} + +func (c *parseContext) match(ch rune) bool { + return c.at(0) == ch +} + +func (c *parseContext) match2(ch rune, ch2 rune) bool { + return c.at(0) == ch && c.at(1) == ch2 +} + +func (c *parseContext) match3(ch rune, ch2 rune, ch3 rune) bool { + return c.at(0) == ch && c.at(1) == ch2 && c.at(2) == ch3 +} + +func (c *parseContext) makeWord(t string, length int, complete bool) *WordType { + rtn := &WordType{Type: t} + rtn.Offset = c.Pos + rtn.QC = c.QC + rtn.Raw = copyRunes(c.Input[c.Pos : c.Pos+length]) + rtn.Complete = complete + c.Pos += length + return rtn +} + +// returns (found, newOffset) +// shell_meta_chars "()<>;&|" +// possible to maybe add ;;& &>> &> |& ;& +func (c *parseContext) parseOp(offset int) (bool, int) { + ch := c.at(offset) + if ch == '(' || ch == ')' || ch == '<' || ch == '>' || ch == ';' || ch == '&' || ch == '|' { + ch2 := c.at(offset + 1) + if ch2 == 0 { + return true, offset + 1 + } + r2 := string([]rune{ch, ch2}) + if r2 == "<<" { + ch3 := c.at(offset + 2) + if ch3 == '-' || ch3 == '<' { + return true, offset + 3 // "<<-" or "<<<" + } + return true, offset + 2 // "<<" + } + if r2 == ">>" || r2 == "&&" || r2 == "||" || r2 == ";;" || r2 == "<<" || r2 == "<&" || r2 == ">&" || r2 == "<>" || r2 == ">|" { + // we don't return '((' here (requires special processing) + return true, offset + 2 + } + return true, offset + 1 + } + return false, 0 +} + +// returns (new-offset, complete) +func (c *parseContext) skipToChar(offset int, endCh rune, allowEsc bool) (int, bool) { + for { + ch := c.at(offset) + if ch == 0 { + return offset, false + } + if allowEsc && ch == '\\' { + if c.at(offset+1) == 0 { + return offset + 1, false + } + offset += 2 + continue + } + if ch == endCh { + return offset + 1, true + } + offset++ + } +} + +// returns (new-offset, complete) +func (c *parseContext) skipToChar2(offset int, endCh rune, endCh2 rune, allowEsc bool) (int, bool) { + for { + ch := c.at(offset) + ch2 := c.at(offset + 1) + if ch == 0 { + return offset, false + } + if ch2 == 0 { + return offset + 1, false + } + if allowEsc && ch == '\\' { + offset += 2 + continue + } + if ch == endCh && ch2 == endCh2 { + return offset + 2, true + } + offset++ + } +} + +func (c *parseContext) parseStrSQ() *WordType { + if !c.match('\'') { + return nil + } + newOffset, complete := c.skipToChar(1, '\'', false) + w := c.makeWord(WordTypeSQ, newOffset, complete) + return w +} + +func (c *parseContext) parseStrDQ() *WordType { + if !c.match('"') { + return nil + } + newContext := c.clone(c.Pos+1, WordTypeDQ) + subWords, eofExit := newContext.tokenizeDQ() + newOffset := newContext.Pos + 1 + w := c.makeWord(WordTypeDQ, newOffset, !eofExit) + w.Subs = subWords + return w +} + +func (c *parseContext) parseStrDDQ() *WordType { + if !c.match2('$', '"') { + return nil + } + newContext := c.clone(c.Pos+2, WordTypeDQ) // use WordTypeDQ (not DDQ) + subWords, eofExit := newContext.tokenizeDQ() + newOffset := newContext.Pos + 2 + w := c.makeWord(WordTypeDDQ, newOffset, !eofExit) + w.Subs = subWords + return w +} + +func (c *parseContext) parseStrBQ() *WordType { + if !c.match('`') { + return nil + } + newContext := c.clone(c.Pos+1, WordTypeBQ) + subWords, eofExit := newContext.tokenizeRaw() + newOffset := newContext.Pos + 1 + w := c.makeWord(WordTypeBQ, newOffset, !eofExit) + w.Subs = subWords + return w +} + +func (c *parseContext) parseStrANSI() *WordType { + if !c.match2('$', '\'') { + return nil + } + newOffset, complete := c.skipToChar(2, '\'', true) + w := c.makeWord(WordTypeDSQ, newOffset, complete) + return w +} + +func (c *parseContext) parseArith(mustComplete bool) *WordType { + if !c.match2('(', '(') { + return nil + } + newOffset, complete := c.skipToChar2(2, ')', ')', false) + if mustComplete && !complete { + return nil + } + w := c.makeWord(WordTypePP, newOffset, complete) + return w +} + +func (c *parseContext) parseExpansion() *WordType { + if !c.match('$') { + return nil + } + if c.match3('$', '(', '(') { + newOffset, complete := c.skipToChar2(3, ')', ')', false) + w := c.makeWord(WordTypeDPP, newOffset, complete) + return w + } + if c.match2('$', '(') { + // subshell + newContext := c.clone(c.Pos+2, WordTypeDP) + subWords, eofExit := newContext.tokenizeRaw() + newOffset := newContext.Pos + 2 + w := c.makeWord(WordTypeDP, newOffset, !eofExit) + w.Subs = subWords + return w + } + if c.match2('$', '[') { + // deprecated arith expansion + newOffset, complete := c.skipToChar(2, ']', false) + w := c.makeWord(WordTypeDB, newOffset, complete) + return w + } + if c.match2('$', '{') { + // variable expansion + newContext := c.clone(c.Pos+2, WordTypeVarBrace) + _, eofExit := newContext.tokenizeVarBrace() + newOffset := newContext.Pos + 2 + w := c.makeWord(WordTypeVarBrace, newOffset, !eofExit) + return w + } + ch2 := c.at(1) + if ch2 == 0 || unicode.IsSpace(ch2) { + // no expansion + return nil + } + newOffset := c.parseSimpleVarName(1) + if newOffset > 1 { + // simple variable name + w := c.makeWord(WordTypeSimpleVar, newOffset, true) + return w + } + if ch2 == '*' || ch2 == '@' || ch2 == '#' || ch2 == '?' || ch2 == '-' || ch2 == '$' || ch2 == '!' || (ch2 >= '0' && ch2 <= '9') { + // single character variable name, e.g. $@, $_, $1, etc. + w := c.makeWord(WordTypeSimpleVar, 2, true) + return w + } + return nil +} + +// returns newOffset +func (c *parseContext) parseSimpleVarName(offset int) int { + first := true + for { + ch := c.at(offset) + if ch == 0 { + return offset + } + if (ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')) || (!first && ch >= '0' && ch <= '9') { + first = false + offset++ + continue + } + return offset + } +} + +func isSimpleVarName(rstr []rune) bool { + if len(rstr) == 0 { + return false + } + for idx, ch := range rstr { + if (ch == '_' || (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z')) || ((idx != 0) && ch >= '0' && ch <= '9') { + continue + } + return false + } + return true +} + +func Tokenize(cmd string) []*WordType { + c := &parseContext{Input: []rune(cmd)} + rtn, _ := c.tokenizeRaw() + return rtn +} diff --git a/wavesrv/pkg/sstore/dbops.go b/wavesrv/pkg/sstore/dbops.go new file mode 100644 index 00000000..abfcbb90 --- /dev/null +++ b/wavesrv/pkg/sstore/dbops.go @@ -0,0 +1,2654 @@ +package sstore + +import ( + "context" + "errors" + "fmt" + "log" + "strconv" + "strings" + "sync" + "time" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/prompt-server/pkg/dbutil" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/google/uuid" + "github.com/jmoiron/sqlx" + "github.com/sawka/txwrap" +) + +const HistoryCols = "h.historyid, h.ts, h.userid, h.sessionid, h.screenid, h.lineid, h.haderror, h.cmdstr, h.remoteownerid, h.remoteid, h.remotename, h.ismetacmd, h.incognito, h.linenum" +const DefaultMaxHistoryItems = 1000 + +var updateWriterCVar = sync.NewCond(&sync.Mutex{}) +var WebScreenPtyPosLock = &sync.Mutex{} +var WebScreenPtyPosDelIntent = make(map[string]bool) // map[screenid + ":" + lineid] -> bool + +type SingleConnDBGetter struct { + SingleConnLock *sync.Mutex +} + +type FeStateType map[string]string + +type TxWrap = txwrap.TxWrap + +var dbWrap *SingleConnDBGetter + +func init() { + dbWrap = &SingleConnDBGetter{SingleConnLock: &sync.Mutex{}} +} + +func (dbg *SingleConnDBGetter) GetDB(ctx context.Context) (*sqlx.DB, error) { + db, err := GetDB(ctx) + if err != nil { + return nil, err + } + dbg.SingleConnLock.Lock() + return db, nil +} + +func (dbg *SingleConnDBGetter) ReleaseDB(db *sqlx.DB) { + dbg.SingleConnLock.Unlock() +} + +func WithTx(ctx context.Context, fn func(tx *TxWrap) error) error { + return txwrap.DBGWithTx(ctx, dbWrap, fn) +} + +func NotifyUpdateWriter() { + // must happen in a goroutine to prevent deadlock. + // update-writer holds this lock while reading from the DB. we can't be holding the DB lock while calling this! + go func() { + updateWriterCVar.L.Lock() + defer updateWriterCVar.L.Unlock() + updateWriterCVar.Signal() + }() +} + +func UpdateWriterCheckMoreData() { + updateWriterCVar.L.Lock() + defer updateWriterCVar.L.Unlock() + for { + updateCount, err := CountScreenUpdates(context.Background()) + if err != nil { + log.Printf("ERROR getting screen update count (sleeping): %v", err) + // will just lead to a Wait() + } + if updateCount > 0 { + break + } + updateWriterCVar.Wait() + } +} + +func NumSessions(ctx context.Context) (int, error) { + var numSessions int + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := "SELECT count(*) FROM session" + numSessions = tx.GetInt(query) + return nil + }) + return numSessions, txErr +} + +func GetAllRemotes(ctx context.Context) ([]*RemoteType, error) { + var rtn []*RemoteType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote ORDER BY remoteidx` + marr := tx.SelectMaps(query) + for _, m := range marr { + rtn = append(rtn, dbutil.FromMap[*RemoteType](m)) + } + return nil + }) + if err != nil { + return nil, err + } + return rtn, nil +} + +func GetRemoteByAlias(ctx context.Context, alias string) (*RemoteType, error) { + var remote *RemoteType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote WHERE remotealias = ?` + m := tx.GetMap(query, alias) + remote = dbutil.FromMap[*RemoteType](m) + return nil + }) + if err != nil { + return nil, err + } + return remote, nil +} + +func GetRemoteById(ctx context.Context, remoteId string) (*RemoteType, error) { + var remote *RemoteType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote WHERE remoteid = ?` + m := tx.GetMap(query, remoteId) + remote = dbutil.FromMap[*RemoteType](m) + return nil + }) + if err != nil { + return nil, err + } + return remote, nil +} + +func GetLocalRemote(ctx context.Context) (*RemoteType, error) { + var remote *RemoteType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote WHERE local` + m := tx.GetMap(query) + remote = dbutil.FromMap[*RemoteType](m) + return nil + }) + if err != nil { + return nil, err + } + return remote, nil +} + +func GetRemoteByCanonicalName(ctx context.Context, cname string) (*RemoteType, error) { + var remote *RemoteType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote WHERE remotecanonicalname = ?` + remote = dbutil.GetMapGen[*RemoteType](tx, query, cname) + return nil + }) + if err != nil { + return nil, err + } + return remote, nil +} + +func UpsertRemote(ctx context.Context, r *RemoteType) error { + if r == nil { + return fmt.Errorf("cannot insert nil remote") + } + if r.RemoteId == "" { + return fmt.Errorf("cannot insert remote without id") + } + if r.RemoteCanonicalName == "" { + return fmt.Errorf("cannot insert remote with canonicalname") + } + if r.RemoteType == "" { + return fmt.Errorf("cannot insert remote without type") + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT remoteid FROM remote WHERE remoteid = ?` + if tx.Exists(query, r.RemoteId) { + tx.Exec(`DELETE FROM remote WHERE remoteid = ?`, r.RemoteId) + } + query = `SELECT remoteid FROM remote WHERE remotecanonicalname = ?` + if tx.Exists(query, r.RemoteCanonicalName) { + return fmt.Errorf("remote has duplicate canonicalname '%s', cannot create", r.RemoteCanonicalName) + } + query = `SELECT remoteid FROM remote WHERE remotealias = ?` + if r.RemoteAlias != "" && tx.Exists(query, r.RemoteAlias) { + return fmt.Errorf("remote has duplicate alias '%s', cannot create", r.RemoteAlias) + } + query = `SELECT COALESCE(max(remoteidx), 0) FROM remote` + maxRemoteIdx := tx.GetInt(query) + r.RemoteIdx = int64(maxRemoteIdx + 1) + query = `INSERT INTO remote + ( remoteid, remotetype, remotealias, remotecanonicalname, remoteuser, remotehost, connectmode, autoinstall, sshopts, remoteopts, lastconnectts, archived, remoteidx, local, statevars, openaiopts) VALUES + (:remoteid,:remotetype,:remotealias,:remotecanonicalname,:remoteuser,:remotehost,:connectmode,:autoinstall,:sshopts,:remoteopts,:lastconnectts,:archived,:remoteidx,:local,:statevars,:openaiopts)` + tx.NamedExec(query, r.ToMap()) + return nil + }) + return txErr +} + +func UpdateRemoteStateVars(ctx context.Context, remoteId string, stateVars map[string]string) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE remote SET statevars = ? WHERE remoteid = ?` + tx.Exec(query, quickJson(stateVars), remoteId) + return nil + }) +} + +func InsertHistoryItem(ctx context.Context, hitem *HistoryItemType) error { + if hitem == nil { + return fmt.Errorf("cannot insert nil history item") + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `INSERT INTO history + ( historyid, ts, userid, sessionid, screenid, lineid, haderror, cmdstr, remoteownerid, remoteid, remotename, ismetacmd, incognito, linenum) VALUES + (:historyid,:ts,:userid,:sessionid,:screenid,:lineid,:haderror,:cmdstr,:remoteownerid,:remoteid,:remotename,:ismetacmd,:incognito,:linenum)` + tx.NamedExec(query, hitem.ToMap()) + return nil + }) + return txErr +} + +func IsIncognitoScreen(ctx context.Context, sessionId string, screenId string) (bool, error) { + return false, nil +} + +const HistoryQueryChunkSize = 1000 + +func _getNextHistoryItem(items []*HistoryItemType, index int, filterFn func(*HistoryItemType) bool) (*HistoryItemType, int) { + for ; index < len(items); index++ { + item := items[index] + if filterFn(item) { + return item, index + } + } + return nil, index +} + +// returns true if done, false if we still need to process more items +func (result *HistoryQueryResult) processItem(item *HistoryItemType, rawOffset int) bool { + if result.prevItems < result.Offset { + result.prevItems++ + return false + } + if len(result.Items) == result.MaxItems { + result.HasMore = true + result.NextRawOffset = rawOffset + return true + } + if len(result.Items) == 0 { + result.RawOffset = rawOffset + } + result.Items = append(result.Items, item) + return false +} + +func runHistoryQueryWithFilter(tx *TxWrap, opts HistoryQueryOpts) (*HistoryQueryResult, error) { + if opts.MaxItems == 0 { + return nil, fmt.Errorf("invalid query, maxitems is 0") + } + rtn := &HistoryQueryResult{Offset: opts.Offset, MaxItems: opts.MaxItems} + var rawOffset int + if opts.RawOffset >= opts.Offset { + rtn.prevItems = opts.Offset + rawOffset = opts.RawOffset + } else { + rawOffset = 0 + } + for { + resultItems, err := runHistoryQuery(tx, opts, rawOffset, HistoryQueryChunkSize) + if err != nil { + return nil, err + } + isDone := false + for resultIdx := 0; resultIdx < len(resultItems); resultIdx++ { + if opts.FilterFn != nil && !opts.FilterFn(resultItems[resultIdx]) { + continue + } + isDone = rtn.processItem(resultItems[resultIdx], rawOffset+resultIdx) + if isDone { + break + } + } + if isDone { + break + } + if len(resultItems) < HistoryQueryChunkSize { + break + } + rawOffset += HistoryQueryChunkSize + } + return rtn, nil +} + +func runHistoryQuery(tx *TxWrap, opts HistoryQueryOpts, realOffset int, itemLimit int) ([]*HistoryItemType, error) { + // check sessionid/screenid format because we are directly inserting them into the SQL + if opts.SessionId != "" { + _, err := uuid.Parse(opts.SessionId) + if err != nil { + return nil, fmt.Errorf("malformed sessionid") + } + } + if opts.ScreenId != "" { + _, err := uuid.Parse(opts.ScreenId) + if err != nil { + return nil, fmt.Errorf("malformed screenid") + } + } + if opts.RemoteId != "" { + _, err := uuid.Parse(opts.RemoteId) + if err != nil { + return nil, fmt.Errorf("malformed remoteid") + } + } + whereClause := "WHERE 1" + var queryArgs []interface{} + hNumStr := "" + if opts.SessionId != "" && opts.ScreenId != "" { + whereClause += fmt.Sprintf(" AND h.sessionid = '%s' AND h.screenid = '%s'", opts.SessionId, opts.ScreenId) + hNumStr = "" + } else if opts.SessionId != "" { + whereClause += fmt.Sprintf(" AND h.sessionid = '%s'", opts.SessionId) + hNumStr = "s" + } else { + hNumStr = "g" + } + if opts.SearchText != "" { + whereClause += " AND h.cmdstr LIKE ? ESCAPE '\\'" + likeArg := opts.SearchText + likeArg = strings.ReplaceAll(likeArg, "%", "\\%") + likeArg = strings.ReplaceAll(likeArg, "_", "\\_") + queryArgs = append(queryArgs, "%"+likeArg+"%") + } + if opts.FromTs > 0 { + whereClause += fmt.Sprintf(" AND h.ts <= %d", opts.FromTs) + } + if opts.RemoteId != "" { + whereClause += fmt.Sprintf(" AND h.remoteid = '%s'", opts.RemoteId) + } + if opts.NoMeta { + whereClause += " AND NOT h.ismetacmd" + } + query := fmt.Sprintf("SELECT %s, ('%s' || CAST((row_number() OVER win) as text)) historynum FROM history h %s WINDOW win AS (ORDER BY h.ts, h.historyid) ORDER BY h.ts DESC, h.historyid DESC LIMIT %d OFFSET %d", HistoryCols, hNumStr, whereClause, itemLimit, realOffset) + marr := tx.SelectMaps(query, queryArgs...) + rtn := make([]*HistoryItemType, len(marr)) + for idx, m := range marr { + hitem := dbutil.FromMap[*HistoryItemType](m) + rtn[idx] = hitem + } + return rtn, nil +} + +func GetHistoryItems(ctx context.Context, opts HistoryQueryOpts) (*HistoryQueryResult, error) { + var rtn *HistoryQueryResult + txErr := WithTx(ctx, func(tx *TxWrap) error { + var err error + rtn, err = runHistoryQueryWithFilter(tx, opts) + if err != nil { + return err + } + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +func GetHistoryItemByLineNum(ctx context.Context, screenId string, lineNum int) (*HistoryItemType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*HistoryItemType, error) { + query := `SELECT * FROM history WHERE screenid = ? AND linenum = ?` + hitem := dbutil.GetMapGen[*HistoryItemType](tx, query, screenId, lineNum) + return hitem, nil + }) +} + +func GetLastHistoryLineNum(ctx context.Context, screenId string) (int, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (int, error) { + query := `SELECT COALESCE(max(linenum), 0) FROM history WHERE screenid = ?` + maxLineNum := tx.GetInt(query, screenId) + return maxLineNum, nil + }) +} + +// includes archived sessions +func GetBareSessions(ctx context.Context) ([]*SessionType, error) { + var rtn []*SessionType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM session ORDER BY archived, sessionidx, archivedts` + tx.Select(&rtn, query) + return nil + }) + if err != nil { + return nil, err + } + return rtn, nil +} + +// does not include archived, finds lowest sessionidx (for resetting active session) +func GetFirstSessionId(ctx context.Context) (string, error) { + var rtn []string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid from session WHERE NOT archived ORDER by sessionidx` + rtn = tx.SelectStrings(query) + return nil + }) + if txErr != nil { + return "", txErr + } + if len(rtn) == 0 { + return "", nil + } + return rtn[0], nil +} + +func GetBareSessionById(ctx context.Context, sessionId string) (*SessionType, error) { + var rtn SessionType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM session WHERE sessionid = ?` + tx.Get(&rtn, query, sessionId) + return nil + }) + if txErr != nil { + return nil, txErr + } + if rtn.SessionId == "" { + return nil, nil + } + return &rtn, nil +} + +func GetAllSessions(ctx context.Context) (*ModelUpdate, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*ModelUpdate, error) { + update := &ModelUpdate{} + query := `SELECT * FROM session ORDER BY archived, sessionidx, archivedts` + tx.Select(&update.Sessions, query) + sessionMap := make(map[string]*SessionType) + for _, session := range update.Sessions { + sessionMap[session.SessionId] = session + session.Full = true + } + query = `SELECT * FROM screen ORDER BY archived, screenidx, archivedts` + update.Screens = dbutil.SelectMapsGen[*ScreenType](tx, query) + for _, screen := range update.Screens { + screen.Full = true + } + query = `SELECT * FROM remote_instance` + riArr := dbutil.SelectMapsGen[*RemoteInstance](tx, query) + for _, ri := range riArr { + s := sessionMap[ri.SessionId] + if s != nil { + s.Remotes = append(s.Remotes, ri) + } + } + query = `SELECT activesessionid FROM client` + update.ActiveSessionId = tx.GetString(query) + return update, nil + }) +} + +func GetScreenLinesById(ctx context.Context, screenId string) (*ScreenLinesType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*ScreenLinesType, error) { + query := `SELECT screenid FROM screen WHERE screenid = ?` + screen := dbutil.GetMappable[*ScreenLinesType](tx, query, screenId) + if screen == nil { + return nil, nil + } + query = `SELECT * FROM line WHERE screenid = ? ORDER BY linenum` + screen.Lines = dbutil.SelectMappable[*LineType](tx, query, screen.ScreenId) + query = `SELECT * FROM cmd WHERE screenid = ?` + screen.Cmds = dbutil.SelectMapsGen[*CmdType](tx, query, screen.ScreenId) + return screen, nil + }) +} + +// includes archived screens +func GetSessionScreens(ctx context.Context, sessionId string) ([]*ScreenType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) ([]*ScreenType, error) { + query := `SELECT * FROM screen WHERE sessionid = ? ORDER BY archived, screenidx, archivedts` + rtn := dbutil.SelectMapsGen[*ScreenType](tx, query, sessionId) + for _, screen := range rtn { + screen.Full = true + } + return rtn, nil + }) +} + +func GetSessionById(ctx context.Context, id string) (*SessionType, error) { + allSessionsUpdate, err := GetAllSessions(ctx) + if err != nil { + return nil, err + } + allSessions := allSessionsUpdate.Sessions + for _, session := range allSessions { + if session.SessionId == id { + return session, nil + } + } + return nil, nil +} + +func GetSessionByName(ctx context.Context, name string) (*SessionType, error) { + var session *SessionType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE name = ?` + sessionId := tx.GetString(query, name) + if sessionId == "" { + return nil + } + var err error + session, err = GetSessionById(tx.Context(), sessionId) + if err != nil { + return err + } + return nil + }) + if txErr != nil { + return nil, txErr + } + return session, nil +} + +// returns sessionId +// if sessionName == "", it will be generated +func InsertSessionWithName(ctx context.Context, sessionName string, activate bool) (*ModelUpdate, error) { + var newScreen *ScreenType + newSessionId := scbase.GenPromptUUID() + txErr := WithTx(ctx, func(tx *TxWrap) error { + names := tx.SelectStrings(`SELECT name FROM session`) + sessionName = fmtUniqueName(sessionName, "workspace-%d", len(names)+1, names) + maxSessionIdx := tx.GetInt(`SELECT COALESCE(max(sessionidx), 0) FROM session`) + query := `INSERT INTO session (sessionid, name, activescreenid, sessionidx, notifynum, archived, archivedts, sharemode) + VALUES (?, ?, '', ?, 0, 0, 0, ?)` + tx.Exec(query, newSessionId, sessionName, maxSessionIdx+1, ShareModeLocal) + screenUpdate, err := InsertScreen(tx.Context(), newSessionId, "", ScreenCreateOpts{}, true) + if err != nil { + return err + } + newScreen = screenUpdate.Screens[0] + if activate { + query = `UPDATE client SET activesessionid = ?` + tx.Exec(query, newSessionId) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + session, err := GetSessionById(ctx, newSessionId) + if err != nil { + return nil, err + } + update := ModelUpdate{ + Sessions: []*SessionType{session}, + Screens: []*ScreenType{newScreen}, + } + if activate { + update.ActiveSessionId = newSessionId + } + return &update, nil +} + +func SetActiveSessionId(ctx context.Context, sessionId string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("cannot switch to session, not found") + } + query = `UPDATE client SET activesessionid = ?` + tx.Exec(query, sessionId) + return nil + }) + return txErr +} + +func GetActiveSessionId(ctx context.Context) (string, error) { + var rtnId string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT activesessionid FROM client` + rtnId = tx.GetString(query) + return nil + }) + return rtnId, txErr +} + +func SetWinSize(ctx context.Context, winSize ClientWinSizeType) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE client SET winsize = ?` + tx.Exec(query, quickJson(winSize)) + return nil + }) + return txErr +} + +func UpdateClientFeOpts(ctx context.Context, feOpts FeOptsType) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE client SET feopts = ?` + tx.Exec(query, quickJson(feOpts)) + return nil + }) + return txErr +} + +func UpdateClientOpenAIOpts(ctx context.Context, aiOpts OpenAIOptsType) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE client SET openaiopts = ?` + tx.Exec(query, quickJson(aiOpts)) + return nil + }) + return txErr +} + +func containsStr(strs []string, testStr string) bool { + for _, s := range strs { + if s == testStr { + return true + } + } + return false +} + +func fmtUniqueName(name string, defaultFmtStr string, startIdx int, strs []string) string { + var fmtStr string + if name != "" { + if !containsStr(strs, name) { + return name + } + fmtStr = name + "-%d" + startIdx = 2 + } else { + fmtStr = defaultFmtStr + } + if strings.Index(fmtStr, "%d") == -1 { + panic("invalid fmtStr: " + fmtStr) + } + for { + testName := fmt.Sprintf(fmtStr, startIdx) + if containsStr(strs, testName) { + startIdx++ + continue + } + return testName + } +} + +func InsertScreen(ctx context.Context, sessionId string, origScreenName string, opts ScreenCreateOpts, activate bool) (*ModelUpdate, error) { + var newScreenId string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ? AND NOT archived` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("cannot create screen, no session found (or session archived)") + } + localRemoteId := tx.GetString(`SELECT remoteid FROM remote WHERE remotealias = ?`, LocalRemoteAlias) + if localRemoteId == "" { + return fmt.Errorf("cannot create screen, no local remote found") + } + maxScreenIdx := tx.GetInt(`SELECT COALESCE(max(screenidx), 0) FROM screen WHERE sessionid = ? AND NOT archived`, sessionId) + var screenName string + if origScreenName == "" { + screenNames := tx.SelectStrings(`SELECT name FROM screen WHERE sessionid = ? AND NOT archived`, sessionId) + screenName = fmtUniqueName("", "s%d", maxScreenIdx+1, screenNames) + } else { + screenName = origScreenName + } + var baseScreen *ScreenType + if opts.HasCopy() { + if opts.BaseScreenId == "" { + return fmt.Errorf("invalid screen create opts, copy option with no base screen specified") + } + var err error + baseScreen, err = GetScreenById(tx.Context(), opts.BaseScreenId) + if err != nil { + return err + } + if baseScreen == nil { + return fmt.Errorf("cannot create screen, base screen not found") + } + } + newScreenId = scbase.GenPromptUUID() + screen := &ScreenType{ + SessionId: sessionId, + ScreenId: newScreenId, + Name: screenName, + ScreenIdx: int64(maxScreenIdx) + 1, + ScreenOpts: ScreenOptsType{}, + OwnerId: "", + ShareMode: ShareModeLocal, + CurRemote: RemotePtrType{RemoteId: localRemoteId}, + NextLineNum: 1, + SelectedLine: 0, + Anchor: ScreenAnchorType{}, + FocusType: ScreenFocusInput, + Archived: false, + ArchivedTs: 0, + } + query = `INSERT INTO screen ( sessionid, screenid, name, screenidx, screenopts, ownerid, sharemode, webshareopts, curremoteownerid, curremoteid, curremotename, nextlinenum, selectedline, anchor, focustype, archived, archivedts) + VALUES (:sessionid,:screenid,:name,:screenidx,:screenopts,:ownerid,:sharemode,:webshareopts,:curremoteownerid,:curremoteid,:curremotename,:nextlinenum,:selectedline,:anchor,:focustype,:archived,:archivedts)` + tx.NamedExec(query, screen.ToMap()) + if activate { + query = `UPDATE session SET activescreenid = ? WHERE sessionid = ?` + tx.Exec(query, newScreenId, sessionId) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + newScreen, err := GetScreenById(ctx, newScreenId) + if err != nil { + return nil, err + } + update := &ModelUpdate{Screens: []*ScreenType{newScreen}} + if activate { + bareSession, err := GetBareSessionById(ctx, sessionId) + if err != nil { + return nil, txErr + } + update.Sessions = []*SessionType{bareSession} + } + return update, nil +} + +func GetScreenById(ctx context.Context, screenId string) (*ScreenType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*ScreenType, error) { + query := `SELECT * FROM screen WHERE screenid = ?` + screen := dbutil.GetMapGen[*ScreenType](tx, query, screenId) + screen.Full = true + return screen, nil + }) +} + +func FindLineIdByArg(ctx context.Context, screenId string, lineArg string) (string, error) { + var lineId string + txErr := WithTx(ctx, func(tx *TxWrap) error { + lineNum, err := strconv.Atoi(lineArg) + if err == nil { + // valid linenum + query := `SELECT lineid FROM line WHERE screenid = ? AND linenum = ?` + lineId = tx.GetString(query, screenId, lineNum) + } else if len(lineArg) == 8 { + // prefix id string match + query := `SELECT lineid FROM line WHERE screenid = ? AND substr(lineid, 1, 8) = ?` + lineId = tx.GetString(query, screenId, lineArg) + } else { + // id match + query := `SELECT lineid FROM line WHERE screenid = ? AND lineid = ?` + lineId = tx.GetString(query, screenId, lineArg) + } + return nil + }) + if txErr != nil { + return "", txErr + } + return lineId, nil +} + +func GetLineCmdByLineId(ctx context.Context, screenId string, lineId string) (*LineType, *CmdType, error) { + return WithTxRtn3(ctx, func(tx *TxWrap) (*LineType, *CmdType, error) { + query := `SELECT * FROM line WHERE screenid = ? AND lineid = ?` + lineVal := dbutil.GetMappable[*LineType](tx, query, screenId, lineId) + if lineVal == nil { + return nil, nil, nil + } + var cmdRtn *CmdType + query = `SELECT * FROM cmd WHERE screenid = ? AND lineid = ?` + cmdRtn = dbutil.GetMapGen[*CmdType](tx, query, screenId, lineId) + return lineVal, cmdRtn, nil + }) +} + +func InsertLine(ctx context.Context, line *LineType, cmd *CmdType) error { + if line == nil { + return fmt.Errorf("line cannot be nil") + } + if line.LineId == "" { + return fmt.Errorf("line must have lineid set") + } + if line.LineNum != 0 { + return fmt.Errorf("line should not hage linenum set") + } + if cmd != nil && cmd.ScreenId == "" { + return fmt.Errorf("cmd should have screenid set") + } + qjs := dbutil.QuickJson(line.LineState) + if len(qjs) > MaxLineStateSize { + return fmt.Errorf("linestate exceeds maxsize, size[%d] max[%d]", len(qjs), MaxLineStateSize) + } + return WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, line.ScreenId) { + return fmt.Errorf("screen not found, cannot insert line[%s]", line.ScreenId) + } + query = `SELECT nextlinenum FROM screen WHERE screenid = ?` + nextLineNum := tx.GetInt(query, line.ScreenId) + line.LineNum = int64(nextLineNum) + query = `INSERT INTO line ( screenid, userid, lineid, ts, linenum, linenumtemp, linelocal, linetype, linestate, text, renderer, ephemeral, contentheight, star, archived) + VALUES (:screenid,:userid,:lineid,:ts,:linenum,:linenumtemp,:linelocal,:linetype,:linestate,:text,:renderer,:ephemeral,:contentheight,:star,:archived)` + tx.NamedExec(query, dbutil.ToDBMap(line, false)) + query = `UPDATE screen SET nextlinenum = ? WHERE screenid = ?` + tx.Exec(query, nextLineNum+1, line.ScreenId) + if cmd != nil { + cmd.OrigTermOpts = cmd.TermOpts + cmdMap := cmd.ToMap() + query = ` +INSERT INTO cmd ( screenid, lineid, remoteownerid, remoteid, remotename, cmdstr, rawcmdstr, festate, statebasehash, statediffhasharr, termopts, origtermopts, status, cmdpid, remotepid, donets, exitcode, durationms, rtnstate, runout, rtnbasehash, rtndiffhasharr) + VALUES (:screenid,:lineid,:remoteownerid,:remoteid,:remotename,:cmdstr,:rawcmdstr,:festate,:statebasehash,:statediffhasharr,:termopts,:origtermopts,:status,:cmdpid,:remotepid,:donets,:exitcode,:durationms,:rtnstate,:runout,:rtnbasehash,:rtndiffhasharr) +` + tx.NamedExec(query, cmdMap) + } + if isWebShare(tx, line.ScreenId) { + insertScreenLineUpdate(tx, line.ScreenId, line.LineId, UpdateType_LineNew) + } + return nil + }) +} + +func GetCmdByScreenId(ctx context.Context, screenId string, lineId string) (*CmdType, error) { + var cmd *CmdType + err := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM cmd WHERE screenid = ? AND lineid = ?` + cmd = dbutil.GetMapGen[*CmdType](tx, query, screenId, lineId) + return nil + }) + if err != nil { + return nil, err + } + return cmd, nil +} + +func UpdateCmdDoneInfo(ctx context.Context, ck base.CommandKey, donePk *packet.CmdDonePacketType, status string) (*ModelUpdate, error) { + if donePk == nil { + return nil, fmt.Errorf("invalid cmddone packet") + } + if ck.IsEmpty() { + return nil, fmt.Errorf("cannot update cmddoneinfo, empty ck") + } + screenId := ck.GetGroupId() + var rtnCmd *CmdType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE cmd SET status = ?, donets = ?, exitcode = ?, durationms = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, status, donePk.Ts, donePk.ExitCode, donePk.DurationMs, screenId, lineIdFromCK(ck)) + var err error + rtnCmd, err = GetCmdByScreenId(tx.Context(), screenId, lineIdFromCK(ck)) + if err != nil { + return err + } + if isWebShare(tx, screenId) { + insertScreenLineUpdate(tx, screenId, lineIdFromCK(ck), UpdateType_CmdExitCode) + insertScreenLineUpdate(tx, screenId, lineIdFromCK(ck), UpdateType_CmdDurationMs) + insertScreenLineUpdate(tx, screenId, lineIdFromCK(ck), UpdateType_CmdStatus) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + if rtnCmd == nil { + return nil, fmt.Errorf("cmd data not found for ck[%s]", ck) + } + return &ModelUpdate{Cmd: rtnCmd}, nil +} + +func UpdateCmdRtnState(ctx context.Context, ck base.CommandKey, statePtr ShellStatePtr) error { + if ck.IsEmpty() { + return fmt.Errorf("cannot update cmdrtnstate, empty ck") + } + screenId := ck.GetGroupId() + lineId := lineIdFromCK(ck) + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE cmd SET rtnbasehash = ?, rtndiffhasharr = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, statePtr.BaseHash, quickJsonArr(statePtr.DiffHashArr), screenId, lineId) + if isWebShare(tx, screenId) { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_CmdRtnState) + } + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +func AppendCmdErrorPk(ctx context.Context, errPk *packet.CmdErrorPacketType) error { + if errPk == nil || errPk.CK.IsEmpty() { + return fmt.Errorf("invalid cmderror packet (no ck)") + } + screenId := errPk.CK.GetGroupId() + return WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE cmd SET runout = json_insert(runout, '$[#]', ?) WHERE screenid = ? AND lineid = ?` + tx.Exec(query, quickJson(errPk), screenId, lineIdFromCK(errPk.CK)) + return nil + }) +} + +func ReInitFocus(ctx context.Context) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE screen SET focustype = 'input'` + tx.Exec(query) + return nil + }) +} + +func HangupAllRunningCmds(ctx context.Context) error { + return WithTx(ctx, func(tx *TxWrap) error { + var cmdPtrs []CmdPtr + query := `SELECT screenid, lineid FROM cmd WHERE status = ?` + tx.Select(&cmdPtrs, query, CmdStatusRunning) + query = `UPDATE cmd SET status = ? WHERE status = ?` + tx.Exec(query, CmdStatusHangup, CmdStatusRunning) + for _, cmdPtr := range cmdPtrs { + if isWebShare(tx, cmdPtr.ScreenId) { + insertScreenLineUpdate(tx, cmdPtr.ScreenId, cmdPtr.LineId, UpdateType_CmdStatus) + } + } + return nil + }) +} + +// TODO send update +func HangupRunningCmdsByRemoteId(ctx context.Context, remoteId string) ([]*ScreenType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) ([]*ScreenType, error) { + var cmdPtrs []CmdPtr + query := `SELECT screenid, lineid FROM cmd WHERE status = ? AND remoteid = ?` + tx.Select(&cmdPtrs, query, CmdStatusRunning, remoteId) + query = `UPDATE cmd SET status = ? WHERE status = ? AND remoteid = ?` + tx.Exec(query, CmdStatusHangup, CmdStatusRunning, remoteId) + var rtn []*ScreenType + for _, cmdPtr := range cmdPtrs { + if isWebShare(tx, cmdPtr.ScreenId) { + insertScreenLineUpdate(tx, cmdPtr.ScreenId, cmdPtr.LineId, UpdateType_CmdStatus) + } + screen, err := UpdateScreenFocusForDoneCmd(tx.Context(), cmdPtr.ScreenId, cmdPtr.LineId) + if err != nil { + return nil, err + } + if screen != nil { + rtn = append(rtn, screen) + } + } + return rtn, nil + }) +} + +// TODO send update +func HangupCmd(ctx context.Context, ck base.CommandKey) (*ScreenType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*ScreenType, error) { + query := `UPDATE cmd SET status = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, CmdStatusHangup, ck.GetGroupId(), lineIdFromCK(ck)) + if isWebShare(tx, ck.GetGroupId()) { + insertScreenLineUpdate(tx, ck.GetGroupId(), lineIdFromCK(ck), UpdateType_CmdStatus) + } + screen, err := UpdateScreenFocusForDoneCmd(tx.Context(), ck.GetGroupId(), lineIdFromCK(ck)) + if err != nil { + return nil, err + } + return screen, nil + }) +} + +func getNextId(ids []string, delId string) string { + if len(ids) == 0 { + return "" + } + if len(ids) == 1 { + if ids[0] == delId { + return "" + } + return ids[0] + } + for idx := 0; idx < len(ids); idx++ { + if ids[idx] == delId { + var rtnIdx int + if idx == len(ids)-1 { + rtnIdx = idx - 1 + } else { + rtnIdx = idx + 1 + } + return ids[rtnIdx] + } + } + return ids[0] +} + +func SwitchScreenById(ctx context.Context, sessionId string, screenId string) (*ModelUpdate, error) { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE sessionid = ? AND screenid = ?` + if !tx.Exists(query, sessionId, screenId) { + return fmt.Errorf("cannot switch to screen, screen=%s does not exist in session=%s", screenId, sessionId) + } + query = `UPDATE session SET activescreenid = ? WHERE sessionid = ?` + tx.Exec(query, screenId, sessionId) + return nil + }) + if txErr != nil { + return nil, txErr + } + bareSession, err := GetBareSessionById(ctx, sessionId) + if err != nil { + return nil, err + } + return &ModelUpdate{ActiveSessionId: sessionId, Sessions: []*SessionType{bareSession}}, nil +} + +// screen may not exist at this point (so don't query screen table) +func cleanScreenCmds(ctx context.Context, screenId string) error { + var removedCmds []string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT lineid FROM cmd WHERE screenid = ? AND lineid NOT IN (SELECT lineid FROM line WHERE screenid = ?)` + removedCmds = tx.SelectStrings(query, screenId, screenId) + query = `DELETE FROM cmd WHERE screenid = ? AND lineid NOT IN (SELECT lineid FROM line WHERE screenid = ?)` + tx.Exec(query, screenId, screenId) + return nil + }) + if txErr != nil { + return txErr + } + for _, lineId := range removedCmds { + DeletePtyOutFile(ctx, screenId, lineId) + } + return nil +} + +func ArchiveScreen(ctx context.Context, sessionId string, screenId string) (UpdatePacket, error) { + var isActive bool + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE sessionid = ? AND screenid = ?` + if !tx.Exists(query, sessionId, screenId) { + return fmt.Errorf("cannot close screen (not found)") + } + if isWebShare(tx, screenId) { + return fmt.Errorf("cannot archive screen while web-sharing. stop web-sharing before trying to archive.") + } + query = `SELECT archived FROM screen WHERE sessionid = ? AND screenid = ?` + closeVal := tx.GetBool(query, sessionId, screenId) + if closeVal { + return nil + } + query = `SELECT count(*) FROM screen WHERE sessionid = ? AND NOT archived` + numScreens := tx.GetInt(query, sessionId) + if numScreens <= 1 { + return fmt.Errorf("cannot archive the last screen in a session") + } + query = `UPDATE screen SET archived = 1, archivedts = ?, screenidx = 0 WHERE sessionid = ? AND screenid = ?` + tx.Exec(query, time.Now().UnixMilli(), sessionId, screenId) + isActive = tx.Exists(`SELECT sessionid FROM session WHERE sessionid = ? AND activescreenid = ?`, sessionId, screenId) + if isActive { + screenIds := tx.SelectStrings(`SELECT screenid FROM screen WHERE sessionid = ? AND NOT archived ORDER BY screenidx`, sessionId) + nextId := getNextId(screenIds, screenId) + tx.Exec(`UPDATE session SET activescreenid = ? WHERE sessionid = ?`, nextId, sessionId) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + newScreen, err := GetScreenById(ctx, screenId) + if err != nil { + return nil, fmt.Errorf("cannot retrive archived screen: %w", err) + } + update := &ModelUpdate{Screens: []*ScreenType{newScreen}} + if isActive { + bareSession, err := GetBareSessionById(ctx, sessionId) + if err != nil { + return nil, err + } + update.Sessions = []*SessionType{bareSession} + } + return update, nil +} + +func UnArchiveScreen(ctx context.Context, sessionId string, screenId string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE sessionid = ? AND screenid = ? AND archived` + if !tx.Exists(query, sessionId, screenId) { + return fmt.Errorf("cannot re-open screen (not found or not archived)") + } + maxScreenIdx := tx.GetInt(`SELECT COALESCE(max(screenidx), 0) FROM screen WHERE sessionid = ? AND NOT archived`, sessionId) + query = `UPDATE screen SET archived = 0, screenidx = ? WHERE sessionid = ? AND screenid = ?` + tx.Exec(query, maxScreenIdx+1, sessionId, screenId) + return nil + }) + return txErr +} + +func PurgeScreen(ctx context.Context, screenId string, sessionDel bool) (UpdatePacket, error) { + var sessionId string + var isActive bool + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, screenId) { + return fmt.Errorf("cannot purge screen (not found)") + } + webSharing := isWebShare(tx, screenId) + if !sessionDel { + query = `SELECT sessionid FROM screen WHERE screenid = ?` + sessionId = tx.GetString(query, screenId) + if sessionId == "" { + return fmt.Errorf("cannot purge screen (no sessionid)") + } + query = `SELECT count(*) FROM screen WHERE sessionid = ? AND NOT archived` + numScreens := tx.GetInt(query, sessionId) + if numScreens <= 1 { + return fmt.Errorf("cannot purge the last screen in a session") + } + isActive = tx.Exists(`SELECT sessionid FROM session WHERE sessionid = ? AND activescreenid = ?`, sessionId, screenId) + if isActive { + screenIds := tx.SelectStrings(`SELECT screenid FROM screen WHERE sessionid = ? AND NOT archived ORDER BY screenidx`, sessionId) + nextId := getNextId(screenIds, screenId) + tx.Exec(`UPDATE session SET activescreenid = ? WHERE sessionid = ?`, nextId, sessionId) + } + } + query = `DELETE FROM screen WHERE screenid = ?` + tx.Exec(query, screenId) + query = `DELETE FROM history WHERE screenid = ?` + tx.Exec(query, screenId) + query = `DELETE FROM line WHERE screenid = ?` + tx.Exec(query, screenId) + query = `DELETE FROM cmd WHERE screenid = ?` + tx.Exec(query, screenId) + if webSharing { + insertScreenDelUpdate(tx, screenId) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + delErr := DeleteScreenDir(ctx, screenId) + if delErr != nil { + log.Printf("error removing screendir") + } + if sessionDel { + return nil, nil + } + update := &ModelUpdate{} + update.Screens = []*ScreenType{&ScreenType{SessionId: sessionId, ScreenId: screenId, Remove: true}} + if isActive { + bareSession, err := GetBareSessionById(ctx, sessionId) + if err != nil { + return nil, err + } + update.Sessions = []*SessionType{bareSession} + } + return update, nil +} + +func GetRemoteState(ctx context.Context, sessionId string, screenId string, remotePtr RemotePtrType) (*packet.ShellState, *ShellStatePtr, error) { + ssptr, err := GetRemoteStatePtr(ctx, sessionId, screenId, remotePtr) + if err != nil { + return nil, nil, err + } + if ssptr == nil { + return nil, nil, nil + } + state, err := GetFullState(ctx, *ssptr) + if err != nil { + return nil, nil, err + } + return state, ssptr, err +} + +func GetRemoteStatePtr(ctx context.Context, sessionId string, screenId string, remotePtr RemotePtrType) (*ShellStatePtr, error) { + var ssptr *ShellStatePtr + txErr := WithTx(ctx, func(tx *TxWrap) error { + ri, err := GetRemoteInstance(tx.Context(), sessionId, screenId, remotePtr) + if err != nil { + return err + } + if ri == nil { + return nil + } + ssptr = &ShellStatePtr{ri.StateBaseHash, ri.StateDiffHashArr} + return nil + }) + if txErr != nil { + return nil, txErr + } + return ssptr, nil +} + +func validateSessionScreen(tx *TxWrap, sessionId string, screenId string) error { + if screenId == "" { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("no session found") + } + return nil + } else { + query := `SELECT screenid FROM screen WHERE sessionid = ? AND screenid = ?` + if !tx.Exists(query, sessionId, screenId) { + return fmt.Errorf("no screen found") + } + return nil + } +} + +func GetRemoteInstance(ctx context.Context, sessionId string, screenId string, remotePtr RemotePtrType) (*RemoteInstance, error) { + if remotePtr.IsSessionScope() { + screenId = "" + } + var ri *RemoteInstance + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote_instance WHERE sessionid = ? AND screenid = ? AND remoteownerid = ? AND remoteid = ? AND name = ?` + ri = dbutil.GetMapGen[*RemoteInstance](tx, query, sessionId, screenId, remotePtr.OwnerId, remotePtr.RemoteId, remotePtr.Name) + return nil + }) + if txErr != nil { + return nil, txErr + } + return ri, nil +} + +// internal function for UpdateRemoteState +func updateRIWithState(ctx context.Context, ri *RemoteInstance, stateBase *packet.ShellState, stateDiff *packet.ShellStateDiff) error { + if stateBase != nil { + ri.StateBaseHash = stateBase.GetHashVal(false) + ri.StateDiffHashArr = nil + err := StoreStateBase(ctx, stateBase) + if err != nil { + return err + } + } else if stateDiff != nil { + ri.StateBaseHash = stateDiff.BaseHash + ri.StateDiffHashArr = append(stateDiff.DiffHashArr, stateDiff.GetHashVal(false)) + err := StoreStateDiff(ctx, stateDiff) + if err != nil { + return err + } + } + return nil +} + +func UpdateRemoteState(ctx context.Context, sessionId string, screenId string, remotePtr RemotePtrType, feState FeStateType, stateBase *packet.ShellState, stateDiff *packet.ShellStateDiff) (*RemoteInstance, error) { + if stateBase == nil && stateDiff == nil { + return nil, fmt.Errorf("UpdateRemoteState, must set state or diff") + } + if stateBase != nil && stateDiff != nil { + return nil, fmt.Errorf("UpdateRemoteState, cannot set state and diff") + } + if remotePtr.IsSessionScope() { + screenId = "" + } + var ri *RemoteInstance + txErr := WithTx(ctx, func(tx *TxWrap) error { + err := validateSessionScreen(tx, sessionId, screenId) + if err != nil { + return fmt.Errorf("cannot update remote instance state: %w", err) + } + query := `SELECT * FROM remote_instance WHERE sessionid = ? AND screenid = ? AND remoteownerid = ? AND remoteid = ? AND name = ?` + ri = dbutil.GetMapGen[*RemoteInstance](tx, query, sessionId, screenId, remotePtr.OwnerId, remotePtr.RemoteId, remotePtr.Name) + if ri == nil { + ri = &RemoteInstance{ + RIId: scbase.GenPromptUUID(), + Name: remotePtr.Name, + SessionId: sessionId, + ScreenId: screenId, + RemoteOwnerId: remotePtr.OwnerId, + RemoteId: remotePtr.RemoteId, + FeState: feState, + } + err = updateRIWithState(tx.Context(), ri, stateBase, stateDiff) + if err != nil { + return err + } + query = `INSERT INTO remote_instance ( riid, name, sessionid, screenid, remoteownerid, remoteid, festate, statebasehash, statediffhasharr) + VALUES (:riid,:name,:sessionid,:screenid,:remoteownerid,:remoteid,:festate,:statebasehash,:statediffhasharr)` + tx.NamedExec(query, ri.ToMap()) + return nil + } else { + query = `UPDATE remote_instance SET festate = ?, statebasehash = ?, statediffhasharr = ? WHERE riid = ?` + ri.FeState = feState + err = updateRIWithState(tx.Context(), ri, stateBase, stateDiff) + if err != nil { + return err + } + tx.Exec(query, quickJson(ri.FeState), ri.StateBaseHash, quickJsonArr(ri.StateDiffHashArr), ri.RIId) + return nil + } + }) + return ri, txErr +} + +func UpdateCurRemote(ctx context.Context, screenId string, remotePtr RemotePtrType) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, screenId) { + return fmt.Errorf("cannot update curremote: no screen found") + } + query = `UPDATE screen SET curremoteownerid = ?, curremoteid = ?, curremotename = ? WHERE screenid = ?` + tx.Exec(query, remotePtr.OwnerId, remotePtr.RemoteId, remotePtr.Name, screenId) + return nil + }) +} + +func reorderStrings(strs []string, toMove string, newIndex int) []string { + if toMove == "" { + return strs + } + var newStrs []string + if newIndex < 0 { + newStrs = append(newStrs, toMove) + } + for _, sval := range strs { + if len(newStrs) == newIndex { + newStrs = append(newStrs, toMove) + } + if sval != toMove { + newStrs = append(newStrs, sval) + } + } + if newIndex >= len(newStrs) { + newStrs = append(newStrs, toMove) + } + return newStrs +} + +func ReIndexSessions(ctx context.Context, sessionId string, newIndex int) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE NOT archived ORDER BY sessionidx, name, sessionid` + ids := tx.SelectStrings(query) + if sessionId != "" { + ids = reorderStrings(ids, sessionId, newIndex) + } + query = `UPDATE session SET sessionid = ? WHERE sessionid = ?` + for idx, id := range ids { + tx.Exec(query, id, idx+1) + } + return nil + }) + return txErr +} + +func SetSessionName(ctx context.Context, sessionId string, name string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("session does not exist") + } + query = `SELECT archived FROM session WHERE sessionid = ?` + isArchived := tx.GetBool(query, sessionId) + if !isArchived { + query = `SELECT sessionid FROM session WHERE name = ? AND NOT archived` + dupSessionId := tx.GetString(query, name) + if dupSessionId == sessionId { + return nil + } + if dupSessionId != "" { + return fmt.Errorf("invalid duplicate session name '%s'", name) + } + } + query = `UPDATE session SET name = ? WHERE sessionid = ?` + tx.Exec(query, name, sessionId) + return nil + }) + return txErr +} + +func SetScreenName(ctx context.Context, sessionId string, screenId string, name string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE sessionid = ? AND screenid = ?` + if !tx.Exists(query, sessionId, screenId) { + return fmt.Errorf("screen does not exist") + } + query = `UPDATE screen SET name = ? WHERE sessionid = ? AND screenid = ?` + tx.Exec(query, name, sessionId, screenId) + return nil + }) + return txErr +} + +func ArchiveScreenLines(ctx context.Context, screenId string) (*ModelUpdate, error) { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, screenId) { + return fmt.Errorf("screen does not exist") + } + fmt.Printf("** archive-screen-lines: %s\n", screenId) + if isWebShare(tx, screenId) { + query = `INSERT INTO screenupdate (screenid, lineid, updatetype, updatets) + SELECT screenid, lineid, ?, ? FROM line WHERE screenid = ? AND archived = 0` + tx.Exec(query, UpdateType_LineDel, time.Now().UnixMilli(), screenId) + NotifyUpdateWriter() + query = `SELECT count(*) FROM line WHERE screenid = ? AND archived = 0` + count := tx.GetInt(query, screenId) + fmt.Printf("** archive-screen-lines: wrote into screenupdate: %d\n", count) + } + query = `UPDATE line SET archived = 1 WHERE screenid = ? AND archived = 0` + tx.Exec(query, screenId) + return nil + }) + if txErr != nil { + return nil, txErr + } + screenLines, err := GetScreenLinesById(ctx, screenId) + if err != nil { + return nil, err + } + return &ModelUpdate{ScreenLines: screenLines}, nil +} + +func PurgeScreenLines(ctx context.Context, screenId string) (*ModelUpdate, error) { + var lineIds []string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT lineid FROM line WHERE screenid = ?` + lineIds = tx.SelectStrings(query, screenId) + query = `DELETE FROM line WHERE screenid = ?` + tx.Exec(query, screenId) + query = `DELETE FROM history WHERE screenid = ?` + tx.Exec(query, screenId) + query = `UPDATE screen SET nextlinenum = 1 WHERE screenid = ?` + tx.Exec(query, screenId) + return nil + }) + if txErr != nil { + return nil, txErr + } + go cleanScreenCmds(context.Background(), screenId) + screen, err := GetScreenById(ctx, screenId) + if err != nil { + return nil, err + } + screenLines, err := GetScreenLinesById(ctx, screenId) + if err != nil { + return nil, err + } + for _, lineId := range lineIds { + line := &LineType{ + ScreenId: screenId, + LineId: lineId, + Remove: true, + } + screenLines.Lines = append(screenLines.Lines, line) + } + return &ModelUpdate{Screens: []*ScreenType{screen}, ScreenLines: screenLines}, nil +} + +func GetRunningScreenCmds(ctx context.Context, screenId string) ([]*CmdType, error) { + var rtn []*CmdType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM cmd WHERE screenid = ? AND status = ?` + rtn = dbutil.SelectMapsGen[*CmdType](tx, query, screenId, CmdStatusRunning) + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +func UpdateCmdTermOpts(ctx context.Context, screenId string, lineId string, termOpts TermOpts) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE cmd SET termopts = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, termOpts, screenId, lineId) + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_CmdTermOpts) + return nil + }) + return txErr +} + +// returns riids of deleted RIs +func ScreenReset(ctx context.Context, screenId string) ([]*RemoteInstance, error) { + return WithTxRtn(ctx, func(tx *TxWrap) ([]*RemoteInstance, error) { + query := `SELECT sessionid FROM screen WHERE screenid = ?` + sessionId := tx.GetString(query, screenId) + if sessionId == "" { + return nil, fmt.Errorf("screen does not exist") + } + query = `SELECT riid FROM remote_instance WHERE sessionid = ? AND screenid = ?` + riids := tx.SelectStrings(query, sessionId, screenId) + var delRis []*RemoteInstance + for _, riid := range riids { + ri := &RemoteInstance{SessionId: sessionId, ScreenId: screenId, RIId: riid, Remove: true} + delRis = append(delRis, ri) + } + query = `DELETE FROM remote_instance WHERE sessionid = ? AND screenid = ?` + tx.Exec(query, sessionId, screenId) + return delRis, nil + }) +} + +func PurgeSession(ctx context.Context, sessionId string) (UpdatePacket, error) { + var newActiveSessionId string + var screenIds []string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("session does not exist") + } + query = `SELECT screenid FROM screen WHERE sessionid = ?` + screenIds = tx.SelectStrings(query, sessionId) + for _, screenId := range screenIds { + _, err := PurgeScreen(tx.Context(), screenId, true) + if err != nil { + return fmt.Errorf("error purging screen[%s]: %v", screenId, err) + } + } + query = `DELETE FROM session WHERE sessionid = ?` + tx.Exec(query, sessionId) + newActiveSessionId, _ = fixActiveSessionId(tx.Context()) + return nil + }) + if txErr != nil { + return nil, txErr + } + update := &ModelUpdate{} + if newActiveSessionId != "" { + update.ActiveSessionId = newActiveSessionId + } + update.Sessions = append(update.Sessions, &SessionType{SessionId: sessionId, Remove: true}) + for _, screenId := range screenIds { + update.Screens = append(update.Screens, &ScreenType{ScreenId: screenId, Remove: true}) + } + return update, nil +} + +func fixActiveSessionId(ctx context.Context) (string, error) { + var newActiveSessionId string + txErr := WithTx(ctx, func(tx *TxWrap) error { + curActiveSessionId := tx.GetString("SELECT activesessionid FROM client") + query := `SELECT sessionid FROM session WHERE sessionid = ? AND NOT archived` + if tx.Exists(query, curActiveSessionId) { + return nil + } + var err error + newActiveSessionId, err = GetFirstSessionId(tx.Context()) + if err != nil { + return err + } + tx.Exec("UPDATE client SET activesessionid = ?", newActiveSessionId) + return nil + }) + if txErr != nil { + return "", txErr + } + return newActiveSessionId, nil +} + +func ArchiveSession(ctx context.Context, sessionId string) (*ModelUpdate, error) { + if sessionId == "" { + return nil, fmt.Errorf("invalid blank sessionid") + } + var newActiveSessionId string + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("session does not exist") + } + query = `SELECT archived FROM session WHERE sessionid = ?` + isArchived := tx.GetBool(query, sessionId) + if isArchived { + return nil + } + query = `UPDATE session SET archived = 1, archivedts = ? WHERE sessionid = ?` + tx.Exec(query, time.Now().UnixMilli(), sessionId) + newActiveSessionId, _ = fixActiveSessionId(tx.Context()) + return nil + }) + if txErr != nil { + return nil, txErr + } + bareSession, _ := GetBareSessionById(ctx, sessionId) + update := &ModelUpdate{} + if bareSession != nil { + update.Sessions = append(update.Sessions, bareSession) + } + if newActiveSessionId != "" { + update.ActiveSessionId = newActiveSessionId + } + return update, nil +} + +func UnArchiveSession(ctx context.Context, sessionId string, activate bool) (*ModelUpdate, error) { + if sessionId == "" { + return nil, fmt.Errorf("invalid blank sessionid") + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("session does not exist") + } + query = `SELECT archived FROM session WHERE sessionid = ?` + isArchived := tx.GetBool(query, sessionId) + if !isArchived { + return nil + } + query = `UPDATE session SET archived = 0, archivedts = 0 WHERE sessionid = ?` + tx.Exec(query, sessionId) + if activate { + query = `UPDATE client SET activesessionid = ?` + tx.Exec(query, sessionId) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + bareSession, _ := GetBareSessionById(ctx, sessionId) + update := &ModelUpdate{} + if bareSession != nil { + update.Sessions = append(update.Sessions, bareSession) + } + if activate { + update.ActiveSessionId = sessionId + } + return update, nil +} + +func GetSessionStats(ctx context.Context, sessionId string) (*SessionStatsType, error) { + rtn := &SessionStatsType{SessionId: sessionId} + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT sessionid FROM session WHERE sessionid = ?` + if !tx.Exists(query, sessionId) { + return fmt.Errorf("not found") + } + query = `SELECT count(*) FROM screen WHERE sessionid = ? AND NOT archived` + rtn.NumScreens = tx.GetInt(query, sessionId) + query = `SELECT count(*) FROM screen WHERE sessionid = ? AND archived` + rtn.NumArchivedScreens = tx.GetInt(query, sessionId) + query = `SELECT count(*) FROM line WHERE screenid IN (SELECT screenid FROM screen WHERE sessionid = ?)` + rtn.NumLines = tx.GetInt(query, sessionId) + query = `SELECT count(*) FROM cmd WHERE screenid IN (SELECT screenid FROM screen WHERE sessionid = ?)` + rtn.NumCmds = tx.GetInt(query, sessionId) + return nil + }) + if txErr != nil { + return nil, txErr + } + diskSize, err := SessionDiskSize(sessionId) + if err != nil { + return nil, err + } + rtn.DiskStats = diskSize + return rtn, nil +} + +const ( + RemoteField_Alias = "alias" // string + RemoteField_ConnectMode = "connectmode" // string + RemoteField_AutoInstall = "autoinstall" // bool + RemoteField_SSHKey = "sshkey" // string + RemoteField_SSHPassword = "sshpassword" // string + RemoteField_Color = "color" // string +) + +// editMap: alias, connectmode, autoinstall, sshkey, color, sshpassword (from constants) +func UpdateRemote(ctx context.Context, remoteId string, editMap map[string]interface{}) (*RemoteType, error) { + var rtn *RemoteType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT remoteid FROM remote WHERE remoteid = ?` + if !tx.Exists(query, remoteId) { + return fmt.Errorf("remote not found") + } + if alias, found := editMap[RemoteField_Alias]; found { + query = `SELECT remoteid FROM remote WHERE remotealias = ? AND remoteid <> ?` + if alias != "" && tx.Exists(query, alias, remoteId) { + return fmt.Errorf("remote has duplicate alias, cannot update") + } + query = `UPDATE remote SET remotealias = ? WHERE remoteid = ?` + tx.Exec(query, alias, remoteId) + } + if mode, found := editMap[RemoteField_ConnectMode]; found { + query = `UPDATE remote SET connectmode = ? WHERE remoteid = ?` + tx.Exec(query, mode, remoteId) + } + if autoInstall, found := editMap[RemoteField_AutoInstall]; found { + query = `UPDATE remote SET autoinstall = ? WHERE remoteid = ?` + tx.Exec(query, autoInstall, remoteId) + } + if sshKey, found := editMap[RemoteField_SSHKey]; found { + query = `UPDATE remote SET sshopts = json_set(sshopts, '$.sshidentity', ?) WHERE remoteid = ?` + tx.Exec(query, sshKey, remoteId) + } + if sshPassword, found := editMap[RemoteField_SSHPassword]; found { + query = `UPDATE remote SET sshopts = json_set(sshopts, '$.sshpassword', ?) WHERE remoteid = ?` + tx.Exec(query, sshPassword, remoteId) + } + if color, found := editMap[RemoteField_Color]; found { + query = `UPDATE remote SET remoteopts = json_set(remoteopts, '$.color', ?) WHERE remoteid = ?` + tx.Exec(query, color, remoteId) + } + var err error + rtn, err = GetRemoteById(tx.Context(), remoteId) + if err != nil { + return err + } + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +const ( + ScreenField_AnchorLine = "anchorline" // int + ScreenField_AnchorOffset = "anchoroffset" // int + ScreenField_SelectedLine = "selectedline" // int + ScreenField_Focus = "focustype" // string + ScreenField_TabColor = "tabcolor" // string + ScreenField_PTerm = "pterm" // string + ScreenField_Name = "name" // string + ScreenField_ShareName = "sharename" // string +) + +func UpdateScreen(ctx context.Context, screenId string, editMap map[string]interface{}) (*ScreenType, error) { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, screenId) { + return fmt.Errorf("screen not found") + } + if anchorLine, found := editMap[ScreenField_AnchorLine]; found { + query = `UPDATE screen SET anchor = json_set(anchor, '$.anchorline', ?) WHERE screenid = ?` + tx.Exec(query, anchorLine, screenId) + } + if anchorOffset, found := editMap[ScreenField_AnchorOffset]; found { + query = `UPDATE screen SET anchor = json_set(anchor, '$.anchoroffset', ?) WHERE screenid = ?` + tx.Exec(query, anchorOffset, screenId) + } + if sline, found := editMap[ScreenField_SelectedLine]; found { + query = `UPDATE screen SET selectedline = ? WHERE screenid = ?` + tx.Exec(query, sline, screenId) + if isWebShare(tx, screenId) { + insertScreenUpdate(tx, screenId, UpdateType_ScreenSelectedLine) + } + } + if focusType, found := editMap[ScreenField_Focus]; found { + query = `UPDATE screen SET focustype = ? WHERE screenid = ?` + tx.Exec(query, focusType, screenId) + } + if tabColor, found := editMap[ScreenField_TabColor]; found { + query = `UPDATE screen SET screenopts = json_set(screenopts, '$.tabcolor', ?) WHERE screenid = ?` + tx.Exec(query, tabColor, screenId) + } + if pterm, found := editMap[ScreenField_PTerm]; found { + query = `UPDATE screen SET screenopts = json_set(screenopts, '$.pterm', ?) WHERE screenid = ?` + tx.Exec(query, pterm, screenId) + } + if name, found := editMap[ScreenField_Name]; found { + query = `UPDATE screen SET name = ? WHERE screenid = ?` + tx.Exec(query, name, screenId) + } + if shareName, found := editMap[ScreenField_ShareName]; found { + if !isWebShare(tx, screenId) { + return fmt.Errorf("cannot set sharename, screen is not web-shared") + } + query = `UPDATE screen SET webshareopts = json_set(webshareopts, '$.sharename', ?) WHERE screenid = ?` + tx.Exec(query, shareName, screenId) + insertScreenUpdate(tx, screenId, UpdateType_ScreenName) + } + return nil + }) + if txErr != nil { + return nil, txErr + } + return GetScreenById(ctx, screenId) +} + +func GetLineResolveItems(ctx context.Context, screenId string) ([]ResolveItem, error) { + var rtn []ResolveItem + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT lineid as id, linenum as num, archived as hidden FROM line WHERE screenid = ? ORDER BY linenum` + tx.Select(&rtn, query, screenId) + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +func UpdateScreenFocusForDoneCmd(ctx context.Context, screenId string, lineId string) (*ScreenType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*ScreenType, error) { + query := `SELECT screenid + FROM screen s + WHERE s.screenid = ? AND s.focustype = ? + AND s.selectedline IN (SELECT linenum FROM line l WHERE l.screenid = s.screenid AND l.lineid = ?) + ` + if !tx.Exists(query, screenId, ScreenFocusCmd, lineId) { + return nil, nil + } + editMap := make(map[string]interface{}) + editMap[ScreenField_Focus] = ScreenFocusInput + screen, err := UpdateScreen(tx.Context(), screenId, editMap) + if err != nil { + return nil, err + } + return screen, nil + }) +} + +func StoreStateBase(ctx context.Context, state *packet.ShellState) error { + stateBase := &StateBase{ + Version: state.Version, + Ts: time.Now().UnixMilli(), + } + stateBase.BaseHash, stateBase.Data = state.EncodeAndHash() + // envMap := shexec.DeclMapFromState(state) + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT basehash FROM state_base WHERE basehash = ?` + if tx.Exists(query, stateBase.BaseHash) { + return nil + } + query = `INSERT INTO state_base (basehash, ts, version, data) VALUES (:basehash,:ts,:version,:data)` + tx.NamedExec(query, stateBase) + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +func StoreStateDiff(ctx context.Context, diff *packet.ShellStateDiff) error { + stateDiff := &StateDiff{ + BaseHash: diff.BaseHash, + Ts: time.Now().UnixMilli(), + DiffHashArr: diff.DiffHashArr, + } + stateDiff.DiffHash, stateDiff.Data = diff.EncodeAndHash() + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT basehash FROM state_base WHERE basehash = ?` + if stateDiff.BaseHash == "" || !tx.Exists(query, stateDiff.BaseHash) { + return fmt.Errorf("cannot store statediff, basehash:%s does not exist", stateDiff.BaseHash) + } + query = `SELECT diffhash FROM state_diff WHERE diffhash = ?` + for idx, diffHash := range stateDiff.DiffHashArr { + if !tx.Exists(query, diffHash) { + return fmt.Errorf("cannot store statediff, diffhash[%d]:%s does not exist", idx, diffHash) + } + } + if tx.Exists(query, stateDiff.DiffHash) { + return nil + } + query = `INSERT INTO state_diff (diffhash, ts, basehash, diffhasharr, data) VALUES (:diffhash,:ts,:basehash,:diffhasharr,:data)` + tx.NamedExec(query, stateDiff.ToMap()) + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +// returns error when not found +func GetFullState(ctx context.Context, ssPtr ShellStatePtr) (*packet.ShellState, error) { + var state *packet.ShellState + if ssPtr.BaseHash == "" { + return nil, fmt.Errorf("invalid empty basehash") + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + var stateBase StateBase + query := `SELECT * FROM state_base WHERE basehash = ?` + found := tx.Get(&stateBase, query, ssPtr.BaseHash) + if !found { + return fmt.Errorf("ShellState %s not found", ssPtr.BaseHash) + } + state = &packet.ShellState{} + err := state.DecodeShellState(stateBase.Data) + if err != nil { + return err + } + for idx, diffHash := range ssPtr.DiffHashArr { + query = `SELECT * FROM state_diff WHERE diffhash = ?` + stateDiff := dbutil.GetMapGen[*StateDiff](tx, query, diffHash) + if stateDiff == nil { + return fmt.Errorf("ShellStateDiff %s not found", diffHash) + } + var ssDiff packet.ShellStateDiff + err = ssDiff.DecodeShellStateDiff(stateDiff.Data) + if err != nil { + return err + } + newState, err := shexec.ApplyShellStateDiff(*state, ssDiff) + if err != nil { + return fmt.Errorf("GetFullState, diff[%d]:%s: %v", idx, diffHash, err) + } + state = &newState + } + return nil + }) + if txErr != nil { + return nil, txErr + } + if state == nil { + return nil, fmt.Errorf("ShellState not found") + } + return state, nil +} + +func UpdateLineStar(ctx context.Context, screenId string, lineId string, starVal int) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE line SET star = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, starVal, screenId, lineId) + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +func UpdateLineHeight(ctx context.Context, screenId string, lineId string, heightVal int) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE line SET contentheight = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, heightVal, screenId, lineId) + if isWebShare(tx, screenId) { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_LineContentHeight) + } + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +func UpdateLineRenderer(ctx context.Context, screenId string, lineId string, renderer string) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE line SET renderer = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, renderer, screenId, lineId) + if isWebShare(tx, screenId) { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_LineRenderer) + } + return nil + }) +} + +func UpdateLineState(ctx context.Context, screenId string, lineId string, lineState map[string]any) error { + qjs := dbutil.QuickJson(lineState) + if len(qjs) > MaxLineStateSize { + return fmt.Errorf("linestate for line[%s:%s] exceeds maxsize, size[%d] max[%d]", screenId, lineId, len(qjs), MaxLineStateSize) + } + return WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE line SET linestate = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, qjs, screenId, lineId) + if isWebShare(tx, screenId) { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_LineState) + } + return nil + }) +} + +// can return nil, nil if line is not found +func GetLineById(ctx context.Context, screenId string, lineId string) (*LineType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*LineType, error) { + query := `SELECT * FROM line WHERE screenid = ? AND lineid = ?` + line := dbutil.GetMappable[*LineType](tx, query, screenId, lineId) + return line, nil + }) +} + +func SetLineArchivedById(ctx context.Context, screenId string, lineId string, archived bool) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE line SET archived = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, archived, screenId, lineId) + if isWebShare(tx, screenId) { + if archived { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_LineDel) + } else { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_LineNew) + } + } + return nil + }) + return txErr +} + +func purgeCmdByScreenId(ctx context.Context, screenId string, lineId string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `DELETE FROM cmd WHERE screenid = ? AND lineid = ?` + tx.Exec(query, screenId, lineId) + if tx.Err != nil { + // short circuit here because we don't want to delete the ptyfile when the tx will be rolled back + return tx.Err + } + return DeletePtyOutFile(tx.Context(), screenId, lineId) + }) + return txErr +} + +func PurgeLinesByIds(ctx context.Context, screenId string, lineIds []string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + isWS := isWebShare(tx, screenId) + for _, lineId := range lineIds { + query := `DELETE FROM line WHERE screenid = ? AND lineid = ?` + tx.Exec(query, screenId, lineId) + query = `DELETE FROM history WHERE screenid = ? AND lineid = ?` + tx.Exec(query, screenId, lineId) + err := purgeCmdByScreenId(tx.Context(), screenId, lineId) + if err != nil { + return err + } + if isWS { + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_LineDel) + } + } + return nil + }) + return txErr +} + +func GetRIsForScreen(ctx context.Context, sessionId string, screenId string) ([]*RemoteInstance, error) { + var rtn []*RemoteInstance + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM remote_instance WHERE sessionid = ? AND (screenid = '' OR screenid = ?)` + rtn = dbutil.SelectMapsGen[*RemoteInstance](tx, query, sessionId, screenId) + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +func GetCurDayStr() string { + now := time.Now() + dayStr := now.Format("2006-01-02") + return dayStr +} + +func UpdateCurrentActivity(ctx context.Context, update ActivityUpdate) error { + now := time.Now() + dayStr := GetCurDayStr() + txErr := WithTx(ctx, func(tx *TxWrap) error { + var tdata TelemetryData + query := `SELECT tdata FROM activity WHERE day = ?` + found := tx.Get(&tdata, query, dayStr) + if !found { + query = `INSERT INTO activity (day, uploaded, tdata, tzname, tzoffset, clientversion, clientarch, buildtime, osrelease) + VALUES (?, 0, ?, ?, ?, ?, ? , ? , ?)` + tzName, tzOffset := now.Zone() + if len(tzName) > MaxTzNameLen { + tzName = tzName[0:MaxTzNameLen] + } + tx.Exec(query, dayStr, tdata, tzName, tzOffset, scbase.PromptVersion, scbase.ClientArch(), scbase.BuildTime, scbase.MacOSRelease()) + } + tdata.NumCommands += update.NumCommands + tdata.FgMinutes += update.FgMinutes + tdata.ActiveMinutes += update.ActiveMinutes + tdata.OpenMinutes += update.OpenMinutes + tdata.ClickShared += update.ClickShared + tdata.HistoryView += update.HistoryView + tdata.BookmarksView += update.BookmarksView + if update.NumConns > 0 { + tdata.NumConns = update.NumConns + } + query = `UPDATE activity + SET tdata = ?, + clientversion = ?, + buildtime = ? + WHERE day = ?` + tx.Exec(query, tdata, scbase.PromptVersion, scbase.BuildTime, dayStr) + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +func GetNonUploadedActivity(ctx context.Context) ([]*ActivityType, error) { + var rtn []*ActivityType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM activity WHERE uploaded = 0 ORDER BY day DESC LIMIT 30` + tx.Select(&rtn, query) + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +// note, will not mark the current day as uploaded +func MarkActivityAsUploaded(ctx context.Context, activityArr []*ActivityType) error { + dayStr := GetCurDayStr() + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE activity SET uploaded = 1 WHERE day = ?` + for _, activity := range activityArr { + if activity.Day == dayStr { + continue + } + tx.Exec(query, activity.Day) + } + return nil + }) + return txErr +} + +func foundInStrArr(strs []string, s string) bool { + for _, sval := range strs { + if s == sval { + return true + } + } + return false +} + +// newPos is 0-indexed +func reorderStrs(strs []string, toMove string, newPos int) []string { + if !foundInStrArr(strs, toMove) { + return strs + } + var added bool + rtn := make([]string, 0, len(strs)) + for _, s := range strs { + if s == toMove { + continue + } + if len(rtn) == newPos { + added = true + rtn = append(rtn, toMove) + } + rtn = append(rtn, s) + } + if !added { + rtn = append(rtn, toMove) + } + return rtn +} + +// newScreenIdx is 1-indexed +func SetScreenIdx(ctx context.Context, sessionId string, screenId string, newScreenIdx int) error { + if newScreenIdx <= 0 { + return fmt.Errorf("invalid screenidx/pos, must be greater than 0") + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE sessionid = ? AND screenid = ? AND NOT archived` + if !tx.Exists(query, sessionId, screenId) { + return fmt.Errorf("invalid screen, not found (or archived)") + } + query = `SELECT screenid FROM screen WHERE sessionid = ? AND NOT archived ORDER BY screenidx` + screens := tx.SelectStrings(query, sessionId) + newScreens := reorderStrs(screens, screenId, newScreenIdx-1) + query = `UPDATE screen SET screenidx = ? WHERE sessionid = ? AND screenid = ?` + for idx, sid := range newScreens { + tx.Exec(query, idx+1, sessionId, sid) + } + return nil + }) + return txErr +} + +func GetDBVersion(ctx context.Context) (int, error) { + var version int + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT version FROM schema_migrations` + version = tx.GetInt(query) + return nil + }) + return version, txErr +} + +type bookmarkOrderType struct { + BookmarkId string + OrderIdx int64 +} + +func GetBookmarks(ctx context.Context, tag string) ([]*BookmarkType, error) { + var bms []*BookmarkType + txErr := WithTx(ctx, func(tx *TxWrap) error { + var query string + if tag == "" { + query = `SELECT * FROM bookmark` + bms = dbutil.SelectMapsGen[*BookmarkType](tx, query) + } else { + query = `SELECT * FROM bookmark WHERE EXISTS (SELECT 1 FROM json_each(tags) WHERE value = ?)` + bms = dbutil.SelectMapsGen[*BookmarkType](tx, query, tag) + } + bmMap := dbutil.MakeGenMap(bms) + var orders []bookmarkOrderType + query = `SELECT bookmarkid, orderidx FROM bookmark_order WHERE tag = ?` + tx.Select(&orders, query, tag) + for _, bmOrder := range orders { + bm := bmMap[bmOrder.BookmarkId] + if bm != nil { + bm.OrderIdx = bmOrder.OrderIdx + } + } + return nil + }) + if txErr != nil { + return nil, txErr + } + return bms, nil +} + +func GetBookmarkById(ctx context.Context, bookmarkId string, tag string) (*BookmarkType, error) { + var rtn *BookmarkType + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT * FROM bookmark WHERE bookmarkid = ?` + rtn = dbutil.GetMapGen[*BookmarkType](tx, query, bookmarkId) + if rtn == nil { + return nil + } + query = `SELECT orderidx FROM bookmark_order WHERE bookmarkid = ? AND tag = ?` + orderIdx := tx.GetInt(query, bookmarkId, tag) + rtn.OrderIdx = int64(orderIdx) + return nil + }) + if txErr != nil { + return nil, txErr + } + return rtn, nil +} + +func GetBookmarkIdByArg(ctx context.Context, bookmarkArg string) (string, error) { + var rtnId string + txErr := WithTx(ctx, func(tx *TxWrap) error { + if len(bookmarkArg) == 8 { + query := `SELECT bookmarkid FROM bookmark WHERE bookmarkid LIKE (? || '%')` + rtnId = tx.GetString(query, bookmarkArg) + return nil + } + query := `SELECT bookmarkid FROM bookmark WHERE bookmarkid = ?` + rtnId = tx.GetString(query, bookmarkArg) + return nil + }) + if txErr != nil { + return "", txErr + } + return rtnId, nil +} + +func GetBookmarkIdsByCmdStr(ctx context.Context, cmdStr string) ([]string, error) { + return WithTxRtn(ctx, func(tx *TxWrap) ([]string, error) { + query := `SELECT bookmarkid FROM bookmark WHERE cmdstr = ?` + bmIds := tx.SelectStrings(query, cmdStr) + return bmIds, nil + }) +} + +// ignores OrderIdx field +func InsertBookmark(ctx context.Context, bm *BookmarkType) error { + if bm == nil || bm.BookmarkId == "" { + return fmt.Errorf("invalid empty bookmark id") + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT bookmarkid FROM bookmark WHERE bookmarkid = ?` + if tx.Exists(query, bm.BookmarkId) { + return fmt.Errorf("bookmarkid already exists") + } + query = `INSERT INTO bookmark ( bookmarkid, createdts, cmdstr, alias, tags, description) + VALUES (:bookmarkid,:createdts,:cmdstr,:alias,:tags,:description)` + tx.NamedExec(query, bm.ToMap()) + for _, tag := range append(bm.Tags, "") { + query = `SELECT COALESCE(max(orderidx), 0) FROM bookmark_order WHERE tag = ?` + maxOrder := tx.GetInt(query, tag) + query = `INSERT INTO bookmark_order (tag, bookmarkid, orderidx) VALUES (?, ?, ?)` + tx.Exec(query, tag, bm.BookmarkId, maxOrder+1) + } + return nil + }) + return txErr +} + +const ( + BookmarkField_Desc = "desc" + BookmarkField_CmdStr = "cmdstr" +) + +func EditBookmark(ctx context.Context, bookmarkId string, editMap map[string]interface{}) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT bookmarkid FROM bookmark WHERE bookmarkid = ?` + if !tx.Exists(query, bookmarkId) { + return fmt.Errorf("bookmark not found") + } + if desc, found := editMap[BookmarkField_Desc]; found { + query = `UPDATE bookmark SET description = ? WHERE bookmarkid = ?` + tx.Exec(query, desc, bookmarkId) + } + if cmdStr, found := editMap[BookmarkField_CmdStr]; found { + query = `UPDATE bookmark SET cmdstr = ? WHERE bookmarkid = ?` + tx.Exec(query, cmdStr, bookmarkId) + } + return nil + }) + return txErr +} + +func fixupBookmarkOrder(tx *TxWrap) { + query := ` +WITH new_order AS ( + SELECT tag, bookmarkid, row_number() OVER (PARTITION BY tag ORDER BY orderidx) AS newidx FROM bookmark_order +) +UPDATE bookmark_order +SET orderidx = new_order.newidx +FROM new_order +WHERE bookmark_order.tag = new_order.tag AND bookmark_order.bookmarkid = new_order.bookmarkid +` + tx.Exec(query) +} + +func DeleteBookmark(ctx context.Context, bookmarkId string) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT bookmarkid FROM bookmark WHERE bookmarkid = ?` + if !tx.Exists(query, bookmarkId) { + return fmt.Errorf("bookmark not found") + } + query = `DELETE FROM bookmark WHERE bookmarkid = ?` + tx.Exec(query, bookmarkId) + query = `DELETE FROM bookmark_order WHERE bookmarkid = ?` + tx.Exec(query, bookmarkId) + fixupBookmarkOrder(tx) + return nil + }) + return txErr +} + +func CreatePlaybook(ctx context.Context, name string) (*PlaybookType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*PlaybookType, error) { + query := `SELECT playbookid FROM playbook WHERE name = ?` + if tx.Exists(query, name) { + return nil, fmt.Errorf("playbook %q already exists", name) + } + rtn := &PlaybookType{} + rtn.PlaybookId = uuid.New().String() + rtn.PlaybookName = name + query = `INSERT INTO playbook ( playbookid, playbookname, description, entryids) + VALUES (:playbookid,:playbookname,:description,:entryids)` + tx.Exec(query, rtn.ToMap()) + return rtn, nil + }) +} + +func selectPlaybook(tx *TxWrap, playbookId string) *PlaybookType { + query := `SELECT * FROM playbook where playbookid = ?` + playbook := dbutil.GetMapGen[*PlaybookType](tx, query, playbookId) + return playbook +} + +func AddPlaybookEntry(ctx context.Context, entry *PlaybookEntry) error { + if entry.EntryId == "" { + return fmt.Errorf("invalid entryid") + } + return WithTx(ctx, func(tx *TxWrap) error { + playbook := selectPlaybook(tx, entry.PlaybookId) + if playbook == nil { + return fmt.Errorf("cannot add entry, playbook does not exist") + } + query := `SELECT entryid FROM playbook_entry WHERE entryid = ?` + if tx.Exists(query, entry.EntryId) { + return fmt.Errorf("cannot add entry, entryid already exists") + } + query = `INSERT INTO playbook_entry ( entryid, playbookid, description, alias, cmdstr, createdts, updatedts) + VALUES (:entryid,:playbookid,:description,:alias,:cmdstr,:createdts,:updatedts)` + tx.Exec(query, entry) + playbook.EntryIds = append(playbook.EntryIds, entry.EntryId) + query = `UPDATE playbook SET entryids = ? WHERE playbookid = ?` + tx.Exec(query, quickJsonArr(playbook.EntryIds), entry.PlaybookId) + return nil + }) +} + +func RemovePlaybookEntry(ctx context.Context, playbookId string, entryId string) error { + return WithTx(ctx, func(tx *TxWrap) error { + playbook := selectPlaybook(tx, playbookId) + if playbook == nil { + return fmt.Errorf("cannot remove playbook entry, playbook does not exist") + } + query := `SELECT entryid FROM playbook_entry WHERE entryid = ?` + if !tx.Exists(query, entryId) { + return fmt.Errorf("cannot remove playbook entry, entry does not exist") + } + query = `DELETE FROM playbook_entry WHERE entryid = ?` + tx.Exec(query, entryId) + playbook.RemoveEntry(entryId) + query = `UPDATE playbook SET entryids = ? WHERE playbookid = ?` + tx.Exec(query, quickJsonArr(playbook.EntryIds), playbookId) + return nil + }) +} + +func GetPlaybookById(ctx context.Context, playbookId string) (*PlaybookType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (*PlaybookType, error) { + rtn := selectPlaybook(tx, playbookId) + if rtn == nil { + return nil, nil + } + query := `SELECT * FROM playbook_entry WHERE playbookid = ?` + tx.Select(&rtn.Entries, query, playbookId) + rtn.OrderEntries() + return rtn, nil + }) +} + +func getLineIdsFromHistoryItems(historyItems []*HistoryItemType) []string { + var rtn []string + for _, hitem := range historyItems { + if hitem.LineId != "" { + rtn = append(rtn, hitem.LineId) + } + } + return rtn +} + +func GetLineCmdsFromHistoryItems(ctx context.Context, historyItems []*HistoryItemType) ([]*LineType, []*CmdType, error) { + if len(historyItems) == 0 { + return nil, nil, nil + } + return WithTxRtn3(ctx, func(tx *TxWrap) ([]*LineType, []*CmdType, error) { + lineIdsJsonArr := quickJsonArr(getLineIdsFromHistoryItems(historyItems)) + query := `SELECT * FROM line WHERE lineid IN (SELECT value FROM json_each(?))` + lineArr := dbutil.SelectMappable[*LineType](tx, query, lineIdsJsonArr) + query = `SELECT * FROM cmd WHERE lineid IN (SELECT value FROM json_each(?))` + cmdArr := dbutil.SelectMapsGen[*CmdType](tx, query, lineIdsJsonArr) + return lineArr, cmdArr, nil + }) +} + +func PurgeHistoryByIds(ctx context.Context, historyIds []string) ([]*HistoryItemType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) ([]*HistoryItemType, error) { + query := `SELECT * FROM history WHERE historyid IN (SELECT value FROM json_each(?))` + rtn := dbutil.SelectMapsGen[*HistoryItemType](tx, query, quickJsonArr(historyIds)) + query = `DELETE FROM history WHERE historyid IN (SELECT value FROM json_each(?))` + tx.Exec(query, quickJsonArr(historyIds)) + for _, hitem := range rtn { + if hitem.LineId != "" { + err := PurgeLinesByIds(tx.Context(), hitem.ScreenId, []string{hitem.LineId}) + if err != nil { + return nil, err + } + } + } + return rtn, nil + }) +} + +func CountScreenWebShares(ctx context.Context) (int, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (int, error) { + query := `SELECT count(*) FROM screen WHERE sharemode = ?` + count := tx.GetInt(query, ShareModeWeb) + return count, nil + }) +} + +func CountScreenLines(ctx context.Context, screenId string) (int, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (int, error) { + query := `SELECT count(*) FROM line WHERE screenid = ? AND NOT archived` + lineCount := tx.GetInt(query, screenId) + return lineCount, nil + }) +} + +func CanScreenWebShare(ctx context.Context, screen *ScreenType) error { + if screen == nil { + return fmt.Errorf("cannot share screen, not found") + } + if screen.ShareMode == ShareModeWeb { + return fmt.Errorf("screen is already shared to web") + } + if screen.ShareMode != ShareModeLocal { + return fmt.Errorf("screen cannot be shared, invalid current share mode %q (must be local)", screen.ShareMode) + } + if screen.Archived { + return fmt.Errorf("screen cannot be shared, must un-archive before sharing") + } + webShareCount, err := CountScreenWebShares(ctx) + if err != nil { + return fmt.Errorf("screen cannot be share: error getting webshare count: %v", err) + } + if webShareCount >= MaxWebShareScreenCount { + go UpdateCurrentActivity(context.Background(), ActivityUpdate{WebShareLimit: 1}) + return fmt.Errorf("screen cannot be shared, limited to a maximum of %d shared screen(s)", MaxWebShareScreenCount) + } + lineCount, err := CountScreenLines(ctx, screen.ScreenId) + if err != nil { + return fmt.Errorf("screen cannot be share: error getting screen line count: %v", err) + } + if lineCount > MaxWebShareLineCount { + go UpdateCurrentActivity(context.Background(), ActivityUpdate{WebShareLimit: 1}) + return fmt.Errorf("screen cannot be shared, limited to a maximum of %d lines", MaxWebShareLineCount) + } + return nil +} + +func ScreenWebShareStart(ctx context.Context, screenId string, shareOpts ScreenWebShareOpts) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, screenId) { + return fmt.Errorf("screen does not exist") + } + shareMode := tx.GetString(`SELECT sharemode FROM screen WHERE screenid = ?`, screenId) + if shareMode == ShareModeWeb { + return fmt.Errorf("screen is already shared to web") + } + if shareMode != ShareModeLocal { + return fmt.Errorf("screen cannot be shared, invalid current share mode %q (must be local)", shareMode) + } + query = `UPDATE screen SET sharemode = ?, webshareopts = ? WHERE screenid = ?` + tx.Exec(query, ShareModeWeb, quickJson(shareOpts), screenId) + insertScreenNewUpdate(tx, screenId) + return nil + }) +} + +func ScreenWebShareStop(ctx context.Context, screenId string) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM screen WHERE screenid = ?` + if !tx.Exists(query, screenId) { + return fmt.Errorf("screen does not exist") + } + shareMode := tx.GetString(`SELECT sharemode FROM screen WHERE screenid = ?`, screenId) + if shareMode != ShareModeWeb { + return fmt.Errorf("screen is not currently shared to the web") + } + query = `UPDATE screen SET sharemode = ?, webshareopts = ? WHERE screenid = ?` + tx.Exec(query, ShareModeLocal, "null", screenId) + handleScreenDelUpdate(tx, screenId) + return nil + }) +} + +func isWebShare(tx *TxWrap, screenId string) bool { + return tx.Exists(`SELECT screenid FROM screen WHERE screenid = ? AND sharemode = ?`, screenId, ShareModeWeb) +} + +func insertScreenUpdate(tx *TxWrap, screenId string, updateType string) { + if screenId == "" { + tx.SetErr(errors.New("invalid screen-update, screenid is empty")) + return + } + nowTs := time.Now().UnixMilli() + query := `INSERT INTO screenupdate (screenid, lineid, updatetype, updatets) VALUES (?, ?, ?, ?)` + tx.Exec(query, screenId, "", updateType, nowTs) + NotifyUpdateWriter() +} + +func insertScreenNewUpdate(tx *TxWrap, screenId string) { + nowTs := time.Now().UnixMilli() + query := `INSERT INTO screenupdate (screenid, lineid, updatetype, updatets) + SELECT screenid, lineid, ?, ? FROM line WHERE screenid = ? AND NOT archived ORDER BY linenum DESC` + tx.Exec(query, UpdateType_LineNew, nowTs, screenId) + query = `INSERT INTO screenupdate (screenid, lineid, updatetype, updatets) + SELECT c.screenid, c.lineid, ?, ? FROM cmd c, line l WHERE c.screenid = ? AND l.lineid = c.lineid AND NOT l.archived ORDER BY l.linenum DESC` + tx.Exec(query, UpdateType_PtyPos, nowTs, screenId) + NotifyUpdateWriter() +} + +func handleScreenDelUpdate(tx *TxWrap, screenId string) { + query := `DELETE FROM screenupdate WHERE screenid = ?` + tx.Exec(query, screenId) + query = `DELETE FROM webptypos WHERE screenid = ?` + tx.Exec(query, screenId) + // don't insert UpdateType_ScreenDel (we already processed it in cmdrunner) +} + +func insertScreenDelUpdate(tx *TxWrap, screenId string) { + handleScreenDelUpdate(tx, screenId) + insertScreenUpdate(tx, screenId, UpdateType_ScreenDel) + // don't insert UpdateType_ScreenDel (we already processed it in cmdrunner) +} + +func insertScreenLineUpdate(tx *TxWrap, screenId string, lineId string, updateType string) { + if screenId == "" { + tx.SetErr(errors.New("invalid screen-update, screenid is empty")) + return + } + if lineId == "" { + tx.SetErr(errors.New("invalid screen-update, lineid is empty")) + return + } + if updateType == UpdateType_LineNew || updateType == UpdateType_LineDel { + query := `DELETE FROM screenupdate WHERE screenid = ? AND lineid = ?` + tx.Exec(query, screenId, lineId) + } + query := `INSERT INTO screenupdate (screenid, lineid, updatetype, updatets) VALUES (?, ?, ?, ?)` + tx.Exec(query, screenId, lineId, updateType, time.Now().UnixMilli()) + if updateType == UpdateType_LineNew { + tx.Exec(query, screenId, lineId, UpdateType_PtyPos, time.Now().UnixMilli()) + } + NotifyUpdateWriter() +} + +func GetScreenUpdates(ctx context.Context, maxNum int) ([]*ScreenUpdateType, error) { + return WithTxRtn(ctx, func(tx *TxWrap) ([]*ScreenUpdateType, error) { + var updates []*ScreenUpdateType + query := `SELECT * FROM screenupdate ORDER BY updateid LIMIT ?` + tx.Select(&updates, query, maxNum) + return updates, nil + }) +} + +func RemoveScreenUpdate(ctx context.Context, updateId int64) error { + if updateId < 0 { + return nil // in-memory updates (not from DB) + } + return WithTx(ctx, func(tx *TxWrap) error { + query := `DELETE FROM screenupdate WHERE updateid = ?` + tx.Exec(query, updateId) + return nil + }) +} + +func CountScreenUpdates(ctx context.Context) (int, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (int, error) { + query := `SELECT count(*) FROM screenupdate` + return tx.GetInt(query), nil + }) +} + +func RemoveScreenUpdates(ctx context.Context, updateIds []int64) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `DELETE FROM screenupdate WHERE updateid IN (SELECT value FROM json_each(?))` + tx.Exec(query, quickJsonArr(updateIds)) + return nil + }) +} + +func MaybeInsertPtyPosUpdate(ctx context.Context, screenId string, lineId string) error { + return WithTx(ctx, func(tx *TxWrap) error { + if !isWebShare(tx, screenId) { + return nil + } + insertScreenLineUpdate(tx, screenId, lineId, UpdateType_PtyPos) + return nil + }) +} + +func GetWebPtyPos(ctx context.Context, screenId string, lineId string) (int64, error) { + return WithTxRtn(ctx, func(tx *TxWrap) (int64, error) { + query := `SELECT ptypos FROM webptypos WHERE screenid = ? AND lineid = ?` + ptyPos := tx.GetInt(query, screenId, lineId) + return int64(ptyPos), nil + }) +} + +func DeleteWebPtyPos(ctx context.Context, screenId string, lineId string) error { + fmt.Printf("del webptypos %s:%s\n", screenId, lineId) + return WithTx(ctx, func(tx *TxWrap) error { + query := `DELETE FROM webptypos WHERE screenid = ? AND lineid = ?` + tx.Exec(query, screenId, lineId) + return nil + }) +} + +func SetWebPtyPos(ctx context.Context, screenId string, lineId string, ptyPos int64) error { + return WithTx(ctx, func(tx *TxWrap) error { + query := `SELECT screenid FROM webptypos WHERE screenid = ? AND lineid = ?` + if tx.Exists(query, screenId, lineId) { + query = `UPDATE webptypos SET ptypos = ? WHERE screenid = ? AND lineid = ?` + tx.Exec(query, ptyPos, screenId, lineId) + } else { + query = `INSERT INTO webptypos (screenid, lineid, ptypos) VALUES (?, ?, ?)` + tx.Exec(query, screenId, lineId, ptyPos) + } + return nil + }) +} diff --git a/wavesrv/pkg/sstore/fileops.go b/wavesrv/pkg/sstore/fileops.go new file mode 100644 index 00000000..953b5673 --- /dev/null +++ b/wavesrv/pkg/sstore/fileops.go @@ -0,0 +1,184 @@ +package sstore + +import ( + "context" + "encoding/base64" + "errors" + "fmt" + "io/fs" + "log" + "os" + "path" + + "github.com/commandlinedev/apishell/pkg/cirfile" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/google/uuid" +) + +func CreateCmdPtyFile(ctx context.Context, screenId string, lineId string, maxSize int64) error { + ptyOutFileName, err := scbase.PtyOutFile(screenId, lineId) + if err != nil { + return err + } + f, err := cirfile.CreateCirFile(ptyOutFileName, maxSize) + if err != nil { + return err + } + return f.Close() +} + +func StatCmdPtyFile(ctx context.Context, screenId string, lineId string) (*cirfile.Stat, error) { + ptyOutFileName, err := scbase.PtyOutFile(screenId, lineId) + if err != nil { + return nil, err + } + return cirfile.StatCirFile(ctx, ptyOutFileName) +} + +func AppendToCmdPtyBlob(ctx context.Context, screenId string, lineId string, data []byte, pos int64) (*PtyDataUpdate, error) { + if screenId == "" { + return nil, fmt.Errorf("cannot append to PtyBlob, screenid is not set") + } + if pos < 0 { + return nil, fmt.Errorf("invalid seek pos '%d' in AppendToCmdPtyBlob", pos) + } + ptyOutFileName, err := scbase.PtyOutFile(screenId, lineId) + if err != nil { + return nil, err + } + f, err := cirfile.OpenCirFile(ptyOutFileName) + if err != nil { + return nil, err + } + defer f.Close() + err = f.WriteAt(ctx, data, pos) + if err != nil { + return nil, err + } + data64 := base64.StdEncoding.EncodeToString(data) + update := &PtyDataUpdate{ + ScreenId: screenId, + LineId: lineId, + PtyPos: pos, + PtyData64: data64, + PtyDataLen: int64(len(data)), + } + err = MaybeInsertPtyPosUpdate(ctx, screenId, lineId) + if err != nil { + // just log + log.Printf("error inserting ptypos update %s/%s: %v\n", screenId, lineId, err) + } + return update, nil +} + +// returns (real-offset, data, err) +func ReadFullPtyOutFile(ctx context.Context, screenId string, lineId string) (int64, []byte, error) { + ptyOutFileName, err := scbase.PtyOutFile(screenId, lineId) + if err != nil { + return 0, nil, err + } + f, err := cirfile.OpenCirFile(ptyOutFileName) + if err != nil { + return 0, nil, err + } + defer f.Close() + return f.ReadAll(ctx) +} + +// returns (real-offset, data, err) +func ReadPtyOutFile(ctx context.Context, screenId string, lineId string, offset int64, maxSize int64) (int64, []byte, error) { + ptyOutFileName, err := scbase.PtyOutFile(screenId, lineId) + if err != nil { + return 0, nil, err + } + f, err := cirfile.OpenCirFile(ptyOutFileName) + if err != nil { + return 0, nil, err + } + defer f.Close() + return f.ReadAtWithMax(ctx, offset, maxSize) +} + +type SessionDiskSizeType struct { + NumFiles int + TotalSize int64 + ErrorCount int + Location string +} + +func directorySize(dirName string) (SessionDiskSizeType, error) { + var rtn SessionDiskSizeType + rtn.Location = dirName + entries, err := os.ReadDir(dirName) + if err != nil { + return rtn, err + } + for _, entry := range entries { + if entry.IsDir() { + rtn.ErrorCount++ + continue + } + finfo, err := entry.Info() + if err != nil { + rtn.ErrorCount++ + continue + } + rtn.NumFiles++ + rtn.TotalSize += finfo.Size() + } + return rtn, nil +} + +func SessionDiskSize(sessionId string) (SessionDiskSizeType, error) { + sessionDir, err := scbase.EnsureSessionDir(sessionId) + if err != nil { + return SessionDiskSizeType{}, err + } + return directorySize(sessionDir) +} + +func FullSessionDiskSize() (map[string]SessionDiskSizeType, error) { + sdir := scbase.GetSessionsDir() + entries, err := os.ReadDir(sdir) + if err != nil { + return nil, err + } + rtn := make(map[string]SessionDiskSizeType) + for _, entry := range entries { + if !entry.IsDir() { + continue + } + name := entry.Name() + _, err = uuid.Parse(name) + if err != nil { + continue + } + diskSize, err := directorySize(path.Join(sdir, name)) + if err != nil { + continue + } + rtn[name] = diskSize + } + return rtn, nil +} + +func DeletePtyOutFile(ctx context.Context, screenId string, lineId string) error { + ptyOutFileName, err := scbase.PtyOutFile(screenId, lineId) + if err != nil { + return err + } + err = os.Remove(ptyOutFileName) + if errors.Is(err, fs.ErrNotExist) { + return nil + } + return err +} + +func DeleteScreenDir(ctx context.Context, screenId string) error { + screenDir, err := scbase.EnsureScreenDir(screenId) + if err != nil { + return fmt.Errorf("error getting screendir: %w", err) + } + log.Printf("remove-all %s\n", screenDir) + return os.RemoveAll(screenDir) +} diff --git a/wavesrv/pkg/sstore/map.go b/wavesrv/pkg/sstore/map.go new file mode 100644 index 00000000..c24a5345 --- /dev/null +++ b/wavesrv/pkg/sstore/map.go @@ -0,0 +1,33 @@ +package sstore + +import ( + "context" +) + +func WithTxRtn[RT any](ctx context.Context, fn func(tx *TxWrap) (RT, error)) (RT, error) { + var rtn RT + txErr := WithTx(ctx, func(tx *TxWrap) error { + temp, err := fn(tx) + if err != nil { + return err + } + rtn = temp + return nil + }) + return rtn, txErr +} + +func WithTxRtn3[RT1 any, RT2 any](ctx context.Context, fn func(tx *TxWrap) (RT1, RT2, error)) (RT1, RT2, error) { + var rtn1 RT1 + var rtn2 RT2 + txErr := WithTx(ctx, func(tx *TxWrap) error { + temp1, temp2, err := fn(tx) + if err != nil { + return err + } + rtn1 = temp1 + rtn2 = temp2 + return nil + }) + return rtn1, rtn2, txErr +} diff --git a/wavesrv/pkg/sstore/migrate.go b/wavesrv/pkg/sstore/migrate.go new file mode 100644 index 00000000..0a26264b --- /dev/null +++ b/wavesrv/pkg/sstore/migrate.go @@ -0,0 +1,213 @@ +package sstore + +import ( + "fmt" + "io" + "log" + "os" + "strconv" + "time" + + sh2db "github.com/commandlinedev/prompt-server/db" + _ "github.com/golang-migrate/migrate/v4/database/sqlite3" + _ "github.com/golang-migrate/migrate/v4/source/file" + "github.com/golang-migrate/migrate/v4/source/iofs" + _ "github.com/mattn/go-sqlite3" + + "github.com/golang-migrate/migrate/v4" +) + +const MaxMigration = 22 +const MigratePrimaryScreenVersion = 9 +const CmdScreenSpecialMigration = 13 +const CmdLineSpecialMigration = 20 + +func MakeMigrate() (*migrate.Migrate, error) { + fsVar, err := iofs.New(sh2db.MigrationFS, "migrations") + if err != nil { + return nil, fmt.Errorf("opening iofs: %w", err) + } + // migrationPathUrl := fmt.Sprintf("file://%s", path.Join(wd, "db", "migrations")) + dbUrl := fmt.Sprintf("sqlite3://%s", GetDBName()) + m, err := migrate.NewWithSourceInstance("iofs", fsVar, dbUrl) + // m, err := migrate.New(migrationPathUrl, dbUrl) + if err != nil { + return nil, fmt.Errorf("making migration db[%s]: %w", GetDBName(), err) + } + return m, nil +} + +func copyFile(srcFile string, dstFile string) error { + if srcFile == dstFile { + return fmt.Errorf("cannot copy %s to itself", srcFile) + } + srcFd, err := os.Open(srcFile) + if err != nil { + return fmt.Errorf("cannot open %s: %v", err) + } + defer srcFd.Close() + dstFd, err := os.OpenFile(dstFile, os.O_RDWR|os.O_CREATE|os.O_TRUNC, 0600) + if err != nil { + return fmt.Errorf("cannot open destination file %s: %v", err) + } + _, err = io.Copy(dstFd, srcFd) + if err != nil { + dstFd.Close() + return fmt.Errorf("error copying file: %v", err) + } + return dstFd.Close() +} + +func MigrateUpStep(m *migrate.Migrate, newVersion uint) error { + startTime := time.Now() + err := m.Migrate(newVersion) + if err != nil { + return err + } + if newVersion == CmdScreenSpecialMigration { + mErr := RunMigration13() + if mErr != nil { + return fmt.Errorf("migrating to v%d: %w", newVersion, mErr) + } + } + if newVersion == CmdLineSpecialMigration { + mErr := RunMigration20() + if mErr != nil { + return fmt.Errorf("migrating to v%d: %w", newVersion, mErr) + } + } + log.Printf("[db] migration v%d, elapsed %v\n", newVersion, time.Since(startTime)) + return nil +} + +func MigrateUp(targetVersion uint) error { + m, err := MakeMigrate() + if err != nil { + return err + } + curVersion, dirty, err := MigrateVersion(m) + if dirty { + return fmt.Errorf("cannot migrate up, database is dirty") + } + if err != nil { + return fmt.Errorf("cannot get current migration version: %v", err) + } + if curVersion >= targetVersion { + return nil + } + log.Printf("[db] migrating from %d to %d\n", curVersion, targetVersion) + log.Printf("[db] backing up database %s to %s\n", DBFileName, DBFileNameBackup) + err = copyFile(GetDBName(), GetDBBackupName()) + if err != nil { + return fmt.Errorf("error creating database backup: %v", err) + } + for newVersion := curVersion + 1; newVersion <= targetVersion; newVersion++ { + err = MigrateUpStep(m, newVersion) + if err != nil { + return fmt.Errorf("during migration v%d: %w", newVersion, err) + } + } + log.Printf("[db] migration done, new version = %d\n", targetVersion) + return nil +} + +// returns curVersion, dirty, error +func MigrateVersion(m *migrate.Migrate) (uint, bool, error) { + if m == nil { + var err error + m, err = MakeMigrate() + if err != nil { + return 0, false, err + } + } + curVersion, dirty, err := m.Version() + if err == migrate.ErrNilVersion { + return 0, false, nil + } + return curVersion, dirty, err +} + +func MigrateDown() error { + m, err := MakeMigrate() + if err != nil { + return err + } + err = m.Down() + if err != nil { + return err + } + return nil +} + +func MigrateGoto(n uint) error { + curVersion, _, _ := MigrateVersion(nil) + if curVersion == n { + return nil + } + if curVersion < n { + return MigrateUp(n) + } + m, err := MakeMigrate() + if err != nil { + return err + } + err = m.Migrate(n) + if err != nil { + return err + } + return nil +} + +func TryMigrateUp() error { + curVersion, _, _ := MigrateVersion(nil) + log.Printf("[db] db version = %d\n", curVersion) + if curVersion >= MaxMigration { + return nil + } + err := MigrateUp(MaxMigration) + if err != nil { + return err + } + return MigratePrintVersion() +} + +func MigratePrintVersion() error { + version, dirty, err := MigrateVersion(nil) + if err != nil { + return fmt.Errorf("error getting db version: %v", err) + } + if dirty { + return fmt.Errorf("error db is dirty, version=%d", version) + } + log.Printf("[db] version=%d\n", version) + return nil +} + +func MigrateCommandOpts(opts []string) error { + var err error + if opts[0] == "--migrate-up" { + fmt.Printf("migrate-up %v\n", GetDBName()) + time.Sleep(3 * time.Second) + err = MigrateUp(MaxMigration) + } else if opts[0] == "--migrate-down" { + fmt.Printf("migrate-down %v\n", GetDBName()) + time.Sleep(3 * time.Second) + err = MigrateDown() + } else if opts[0] == "--migrate-goto" { + n, err := strconv.Atoi(opts[1]) + if err == nil { + fmt.Printf("migrate-goto %v => %d\n", GetDBName(), n) + time.Sleep(3 * time.Second) + err = MigrateGoto(uint(n)) + } + } else { + err = fmt.Errorf("invalid migration command") + } + if err != nil && err.Error() == migrate.ErrNoChange.Error() { + err = nil + } + if err != nil { + return err + } + return MigratePrintVersion() +} diff --git a/wavesrv/pkg/sstore/quick.go b/wavesrv/pkg/sstore/quick.go new file mode 100644 index 00000000..c19ce8a0 --- /dev/null +++ b/wavesrv/pkg/sstore/quick.go @@ -0,0 +1,19 @@ +package sstore + +import ( + "github.com/commandlinedev/prompt-server/pkg/dbutil" +) + +var quickSetStr = dbutil.QuickSetStr +var quickSetInt64 = dbutil.QuickSetInt64 +var quickSetInt = dbutil.QuickSetInt +var quickSetBool = dbutil.QuickSetBool +var quickSetBytes = dbutil.QuickSetBytes +var quickSetJson = dbutil.QuickSetJson +var quickSetNullableJson = dbutil.QuickSetNullableJson +var quickSetJsonArr = dbutil.QuickSetJsonArr +var quickNullableJson = dbutil.QuickNullableJson +var quickJson = dbutil.QuickJson +var quickJsonArr = dbutil.QuickJsonArr +var quickScanJson = dbutil.QuickScanJson +var quickValueJson = dbutil.QuickValueJson diff --git a/wavesrv/pkg/sstore/sstore.go b/wavesrv/pkg/sstore/sstore.go new file mode 100644 index 00000000..cd77c255 --- /dev/null +++ b/wavesrv/pkg/sstore/sstore.go @@ -0,0 +1,1295 @@ +package sstore + +import ( + "context" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rand" + "crypto/x509" + "database/sql/driver" + "fmt" + "log" + "os" + "os/user" + "path" + "regexp" + "strings" + "sync" + "time" + + "github.com/commandlinedev/apishell/pkg/base" + "github.com/commandlinedev/apishell/pkg/packet" + "github.com/commandlinedev/apishell/pkg/shexec" + "github.com/commandlinedev/prompt-server/pkg/dbutil" + "github.com/commandlinedev/prompt-server/pkg/scbase" + "github.com/google/uuid" + "github.com/jmoiron/sqlx" + "github.com/sawka/txwrap" + + _ "github.com/mattn/go-sqlite3" +) + +const LineNoHeight = -1 +const DBFileName = "prompt.db" +const DBFileNameBackup = "backup.prompt.db" +const MaxWebShareLineCount = 50 +const MaxWebShareScreenCount = 3 +const MaxLineStateSize = 4 * 1024 // 4k for now, can raise if needed + +const DefaultSessionName = "default" +const LocalRemoteAlias = "local" + +const DefaultCwd = "~" +const APITokenSentinel = "--apitoken--" + +const ( + LineTypeCmd = "cmd" + LineTypeText = "text" + LineTypeOpenAI = "openai" +) + +const ( + LineState_Source = "prompt:source" + LineState_File = "prompt:file" + LineState_Template = "template" + LineState_Mode = "mode" + LineState_Lang = "lang" +) + +const ( + MainViewSession = "session" + MainViewBookmarks = "bookmarks" + MainViewHistory = "history" +) + +const ( + CmdStatusRunning = "running" + CmdStatusDetached = "detached" + CmdStatusError = "error" + CmdStatusDone = "done" + CmdStatusHangup = "hangup" +) + +const ( + CmdRendererOpenAI = "openai" +) + +const ( + OpenAIRoleSystem = "system" + OpenAIRoleUser = "user" + OpenAIRoleAssistant = "assistant" +) + +const ( + RemoteAuthTypeNone = "none" + RemoteAuthTypePassword = "password" + RemoteAuthTypeKey = "key" + RemoteAuthTypeKeyPassword = "key+password" +) + +const ( + ShareModeLocal = "local" + ShareModeWeb = "web" +) + +const ( + ConnectModeStartup = "startup" + ConnectModeAuto = "auto" + ConnectModeManual = "manual" +) + +const ( + RemoteTypeSsh = "ssh" + RemoteTypeOpenAI = "openai" +) + +const ( + ScreenFocusInput = "input" + ScreenFocusCmd = "cmd" +) + +const ( + CmdStoreTypeSession = "session" + CmdStoreTypeScreen = "screen" +) + +const ( + UpdateType_ScreenNew = "screen:new" + UpdateType_ScreenDel = "screen:del" + UpdateType_ScreenSelectedLine = "screen:selectedline" + UpdateType_ScreenName = "screen:sharename" + UpdateType_LineNew = "line:new" + UpdateType_LineDel = "line:del" + UpdateType_LineRenderer = "line:renderer" + UpdateType_LineContentHeight = "line:contentheight" + UpdateType_LineState = "line:state" + UpdateType_CmdStatus = "cmd:status" + UpdateType_CmdTermOpts = "cmd:termopts" + UpdateType_CmdExitCode = "cmd:exitcode" + UpdateType_CmdDurationMs = "cmd:durationms" + UpdateType_CmdRtnState = "cmd:rtnstate" + UpdateType_PtyPos = "pty:pos" +) + +const MaxTzNameLen = 50 + +var globalDBLock = &sync.Mutex{} +var globalDB *sqlx.DB +var globalDBErr error + +func lineIdFromCK(ck base.CommandKey) string { + return ck.GetCmdId() +} + +func GetDBName() string { + scHome := scbase.GetPromptHomeDir() + return path.Join(scHome, DBFileName) +} + +func GetDBBackupName() string { + scHome := scbase.GetPromptHomeDir() + return path.Join(scHome, DBFileNameBackup) +} + +func IsValidConnectMode(mode string) bool { + return mode == ConnectModeStartup || mode == ConnectModeAuto || mode == ConnectModeManual +} + +func GetDB(ctx context.Context) (*sqlx.DB, error) { + if txwrap.IsTxWrapContext(ctx) { + return nil, fmt.Errorf("cannot call GetDB from within a running transaction") + } + globalDBLock.Lock() + defer globalDBLock.Unlock() + if globalDB == nil && globalDBErr == nil { + dbName := GetDBName() + globalDB, globalDBErr = sqlx.Open("sqlite3", fmt.Sprintf("file:%s?cache=shared&mode=rwc&_journal_mode=WAL&_busy_timeout=5000", dbName)) + if globalDBErr != nil { + globalDBErr = fmt.Errorf("opening db[%s]: %w", dbName, globalDBErr) + log.Printf("[db] error: %v\n", globalDBErr) + } else { + log.Printf("[db] successfully opened db %s\n", dbName) + } + } + return globalDB, globalDBErr +} + +func CloseDB() { + globalDBLock.Lock() + defer globalDBLock.Unlock() + if globalDB == nil { + return + } + err := globalDB.Close() + if err != nil { + log.Printf("[db] error closing database: %v\n", err) + } + globalDB = nil +} + +type CmdPtr struct { + ScreenId string + LineId string +} + +type ClientWinSizeType struct { + Width int `json:"width"` + Height int `json:"height"` + Top int `json:"top"` + Left int `json:"left"` + FullScreen bool `json:"fullscreen,omitempty"` +} + +type ActivityUpdate struct { + FgMinutes int + ActiveMinutes int + OpenMinutes int + NumCommands int + ClickShared int + HistoryView int + BookmarksView int + NumConns int + WebShareLimit int + BuildTime string +} + +type ActivityType struct { + Day string `json:"day"` + Uploaded bool `json:"-"` + TData TelemetryData `json:"tdata"` + TzName string `json:"tzname"` + TzOffset int `json:"tzoffset"` + ClientVersion string `json:"clientversion"` + ClientArch string `json:"clientarch"` + BuildTime string `json:"buildtime"` + OSRelease string `json:"osrelease"` +} + +type TelemetryData struct { + NumCommands int `json:"numcommands"` + ActiveMinutes int `json:"activeminutes"` + FgMinutes int `json:"fgminutes"` + OpenMinutes int `json:"openminutes"` + ClickShared int `json:"clickshared,omitempty"` + HistoryView int `json:"historyview,omitempty"` + BookmarksView int `json:"bookmarksview,omitempty"` + NumConns int `json:"numconns"` + WebShareLimit int `json:"websharelimit,omitempty"` +} + +func (tdata TelemetryData) Value() (driver.Value, error) { + return quickValueJson(tdata) +} + +func (tdata *TelemetryData) Scan(val interface{}) error { + return quickScanJson(tdata, val) +} + +type ClientOptsType struct { + NoTelemetry bool `json:"notelemetry,omitempty"` + AcceptedTos int64 `json:"acceptedtos,omitempty"` +} + +type FeOptsType struct { + TermFontSize int `json:"termfontsize,omitempty"` +} + +type ClientData struct { + ClientId string `json:"clientid"` + UserId string `json:"userid"` + UserPrivateKeyBytes []byte `json:"-"` + UserPublicKeyBytes []byte `json:"-"` + UserPrivateKey *ecdsa.PrivateKey `json:"-" dbmap:"-"` + UserPublicKey *ecdsa.PublicKey `json:"-" dbmap:"-"` + ActiveSessionId string `json:"activesessionid"` + WinSize ClientWinSizeType `json:"winsize"` + ClientOpts ClientOptsType `json:"clientopts"` + FeOpts FeOptsType `json:"feopts"` + CmdStoreType string `json:"cmdstoretype"` + DBVersion int `json:"dbversion" dbmap:"-"` + OpenAIOpts *OpenAIOptsType `json:"openaiopts,omitempty" dbmap:"openaiopts"` +} + +func (ClientData) UseDBMap() {} + +func (cdata *ClientData) Clean() *ClientData { + if cdata == nil { + return nil + } + rtn := *cdata + if rtn.OpenAIOpts != nil { + rtn.OpenAIOpts = &OpenAIOptsType{ + Model: cdata.OpenAIOpts.Model, + MaxTokens: cdata.OpenAIOpts.MaxTokens, + MaxChoices: cdata.OpenAIOpts.MaxChoices, + // omit API Token + } + if cdata.OpenAIOpts.APIToken != "" { + rtn.OpenAIOpts.APIToken = APITokenSentinel + } + } + return &rtn +} + +type SessionType struct { + SessionId string `json:"sessionid"` + Name string `json:"name"` + SessionIdx int64 `json:"sessionidx"` + ActiveScreenId string `json:"activescreenid"` + ShareMode string `json:"sharemode"` + NotifyNum int64 `json:"notifynum"` + Archived bool `json:"archived,omitempty"` + ArchivedTs int64 `json:"archivedts,omitempty"` + Remotes []*RemoteInstance `json:"remotes"` + + // only for updates + Remove bool `json:"remove,omitempty"` + Full bool `json:"full,omitempty"` +} + +type SessionStatsType struct { + SessionId string `json:"sessionid"` + NumScreens int `json:"numscreens"` + NumArchivedScreens int `json:"numarchivedscreens"` + NumLines int `json:"numlines"` + NumCmds int `json:"numcmds"` + DiskStats SessionDiskSizeType `json:"diskstats"` +} + +var RemoteNameRe = regexp.MustCompile("^\\*?[a-zA-Z0-9_-]+$") + +type RemotePtrType struct { + OwnerId string `json:"ownerid"` + RemoteId string `json:"remoteid"` + Name string `json:"name"` +} + +func (r RemotePtrType) IsSessionScope() bool { + return strings.HasPrefix(r.Name, "*") +} + +func (rptr *RemotePtrType) GetDisplayName(baseDisplayName string) string { + name := baseDisplayName + if rptr == nil { + return name + } + if rptr.Name != "" { + name = name + ":" + rptr.Name + } + if rptr.OwnerId != "" { + name = "@" + rptr.OwnerId + ":" + name + } + return name +} + +func (r RemotePtrType) Validate() error { + if r.OwnerId != "" { + if _, err := uuid.Parse(r.OwnerId); err != nil { + return fmt.Errorf("invalid ownerid format: %v", err) + } + } + if r.RemoteId != "" { + if _, err := uuid.Parse(r.RemoteId); err != nil { + return fmt.Errorf("invalid remoteid format: %v", err) + } + } + if r.Name != "" { + ok := RemoteNameRe.MatchString(r.Name) + if !ok { + return fmt.Errorf("invalid remote name") + } + } + return nil +} + +func (r RemotePtrType) MakeFullRemoteRef() string { + if r.RemoteId == "" { + return "" + } + if r.OwnerId == "" && r.Name == "" { + return r.RemoteId + } + if r.OwnerId != "" && r.Name == "" { + return fmt.Sprintf("@%s:%s", r.OwnerId, r.RemoteId) + } + if r.OwnerId == "" && r.Name != "" { + return fmt.Sprintf("%s:%s", r.RemoteId, r.Name) + } + return fmt.Sprintf("@%s:%s:%s", r.OwnerId, r.RemoteId, r.Name) +} + +func (h *HistoryItemType) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["historyid"] = h.HistoryId + rtn["ts"] = h.Ts + rtn["userid"] = h.UserId + rtn["sessionid"] = h.SessionId + rtn["screenid"] = h.ScreenId + rtn["lineid"] = h.LineId + rtn["linenum"] = h.LineNum + rtn["haderror"] = h.HadError + rtn["cmdstr"] = h.CmdStr + rtn["remoteownerid"] = h.Remote.OwnerId + rtn["remoteid"] = h.Remote.RemoteId + rtn["remotename"] = h.Remote.Name + rtn["ismetacmd"] = h.IsMetaCmd + rtn["incognito"] = h.Incognito + return rtn +} + +func (h *HistoryItemType) FromMap(m map[string]interface{}) bool { + quickSetStr(&h.HistoryId, m, "historyid") + quickSetInt64(&h.Ts, m, "ts") + quickSetStr(&h.UserId, m, "userid") + quickSetStr(&h.SessionId, m, "sessionid") + quickSetStr(&h.ScreenId, m, "screenid") + quickSetStr(&h.LineId, m, "lineid") + quickSetBool(&h.HadError, m, "haderror") + quickSetStr(&h.CmdStr, m, "cmdstr") + quickSetStr(&h.Remote.OwnerId, m, "remoteownerid") + quickSetStr(&h.Remote.RemoteId, m, "remoteid") + quickSetStr(&h.Remote.Name, m, "remotename") + quickSetBool(&h.IsMetaCmd, m, "ismetacmd") + quickSetStr(&h.HistoryNum, m, "historynum") + quickSetInt64(&h.LineNum, m, "linenum") + quickSetBool(&h.Incognito, m, "incognito") + return true +} + +type ScreenOptsType struct { + TabColor string `json:"tabcolor,omitempty"` + PTerm string `json:"pterm,omitempty"` +} + +type ScreenLinesType struct { + ScreenId string `json:"screenid"` + Lines []*LineType `json:"lines" dbmap:"-"` + Cmds []*CmdType `json:"cmds" dbmap:"-"` +} + +func (ScreenLinesType) UseDBMap() {} + +type ScreenWebShareOpts struct { + ShareName string `json:"sharename"` + ViewKey string `json:"viewkey"` +} + +type ScreenCreateOpts struct { + BaseScreenId string + CopyRemote bool + CopyCwd bool + CopyEnv bool +} + +func (sco ScreenCreateOpts) HasCopy() bool { + return sco.CopyRemote || sco.CopyCwd || sco.CopyEnv +} + +type ScreenType struct { + SessionId string `json:"sessionid"` + ScreenId string `json:"screenid"` + Name string `json:"name"` + ScreenIdx int64 `json:"screenidx"` + ScreenOpts ScreenOptsType `json:"screenopts"` + OwnerId string `json:"ownerid"` + ShareMode string `json:"sharemode"` + WebShareOpts *ScreenWebShareOpts `json:"webshareopts,omitempty"` + CurRemote RemotePtrType `json:"curremote"` + NextLineNum int64 `json:"nextlinenum"` + SelectedLine int64 `json:"selectedline"` + Anchor ScreenAnchorType `json:"anchor"` + FocusType string `json:"focustype"` + Archived bool `json:"archived,omitempty"` + ArchivedTs int64 `json:"archivedts,omitempty"` + + // only for updates + Full bool `json:"full,omitempty"` + Remove bool `json:"remove,omitempty"` +} + +func (s *ScreenType) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["sessionid"] = s.SessionId + rtn["screenid"] = s.ScreenId + rtn["name"] = s.Name + rtn["screenidx"] = s.ScreenIdx + rtn["screenopts"] = quickJson(s.ScreenOpts) + rtn["ownerid"] = s.OwnerId + rtn["sharemode"] = s.ShareMode + rtn["webshareopts"] = quickNullableJson(s.WebShareOpts) + rtn["curremoteownerid"] = s.CurRemote.OwnerId + rtn["curremoteid"] = s.CurRemote.RemoteId + rtn["curremotename"] = s.CurRemote.Name + rtn["nextlinenum"] = s.NextLineNum + rtn["selectedline"] = s.SelectedLine + rtn["anchor"] = quickJson(s.Anchor) + rtn["focustype"] = s.FocusType + rtn["archived"] = s.Archived + rtn["archivedts"] = s.ArchivedTs + return rtn +} + +func (s *ScreenType) FromMap(m map[string]interface{}) bool { + quickSetStr(&s.SessionId, m, "sessionid") + quickSetStr(&s.ScreenId, m, "screenid") + quickSetStr(&s.Name, m, "name") + quickSetInt64(&s.ScreenIdx, m, "screenidx") + quickSetJson(&s.ScreenOpts, m, "screenopts") + quickSetStr(&s.OwnerId, m, "ownerid") + quickSetStr(&s.ShareMode, m, "sharemode") + quickSetNullableJson(&s.WebShareOpts, m, "webshareopts") + quickSetStr(&s.CurRemote.OwnerId, m, "curremoteownerid") + quickSetStr(&s.CurRemote.RemoteId, m, "curremoteid") + quickSetStr(&s.CurRemote.Name, m, "curremotename") + quickSetInt64(&s.NextLineNum, m, "nextlinenum") + quickSetInt64(&s.SelectedLine, m, "selectedline") + quickSetJson(&s.Anchor, m, "anchor") + quickSetStr(&s.FocusType, m, "focustype") + quickSetBool(&s.Archived, m, "archived") + quickSetInt64(&s.ArchivedTs, m, "archivedts") + return true +} + +const ( + LayoutFull = "full" +) + +type LayoutType struct { + Type string `json:"type"` + Parent string `json:"parent,omitempty"` + ZIndex int64 `json:"zindex,omitempty"` + Float bool `json:"float,omitempty"` + Top string `json:"top,omitempty"` + Bottom string `json:"bottom,omitempty"` + Left string `json:"left,omitempty"` + Right string `json:"right,omitempty"` + Width string `json:"width,omitempty"` + Height string `json:"height,omitempty"` +} + +func (l *LayoutType) Scan(val interface{}) error { + return quickScanJson(l, val) +} + +func (l LayoutType) Value() (driver.Value, error) { + return quickValueJson(l) +} + +type ScreenAnchorType struct { + AnchorLine int `json:"anchorline,omitempty"` + AnchorOffset int `json:"anchoroffset,omitempty"` +} + +type HistoryItemType struct { + HistoryId string `json:"historyid"` + Ts int64 `json:"ts"` + UserId string `json:"userid"` + SessionId string `json:"sessionid"` + ScreenId string `json:"screenid"` + LineId string `json:"lineid"` + HadError bool `json:"haderror"` + CmdStr string `json:"cmdstr"` + Remote RemotePtrType `json:"remote"` + IsMetaCmd bool `json:"ismetacmd"` + Incognito bool `json:"incognito,omitempty"` + + // only for updates + Remove bool `json:"remove"` + + // transient (string because of different history orderings) + HistoryNum string `json:"historynum"` + LineNum int64 `json:"linenum"` +} + +type HistoryQueryOpts struct { + Offset int + MaxItems int + FromTs int64 + SearchText string + SessionId string + RemoteId string + ScreenId string + NoMeta bool + RawOffset int + FilterFn func(*HistoryItemType) bool +} + +type HistoryQueryResult struct { + MaxItems int + Items []*HistoryItemType + Offset int // the offset shown to user + RawOffset int // internal offset + HasMore bool + NextRawOffset int // internal offset used by pager for next query + + prevItems int // holds number of items skipped by RawOffset +} + +type TermOpts struct { + Rows int64 `json:"rows"` + Cols int64 `json:"cols"` + FlexRows bool `json:"flexrows,omitempty"` + MaxPtySize int64 `json:"maxptysize,omitempty"` +} + +func (opts *TermOpts) Scan(val interface{}) error { + return quickScanJson(opts, val) +} + +func (opts TermOpts) Value() (driver.Value, error) { + return quickValueJson(opts) +} + +type ShellStatePtr struct { + BaseHash string + DiffHashArr []string +} + +func (ssptr *ShellStatePtr) IsEmpty() bool { + if ssptr == nil || ssptr.BaseHash == "" { + return true + } + return false +} + +type RemoteInstance struct { + RIId string `json:"riid"` + Name string `json:"name"` + SessionId string `json:"sessionid"` + ScreenId string `json:"screenid"` + RemoteOwnerId string `json:"remoteownerid"` + RemoteId string `json:"remoteid"` + FeState map[string]string `json:"festate"` + StateBaseHash string `json:"-"` + StateDiffHashArr []string `json:"-"` + + // only for updates + Remove bool `json:"remove,omitempty"` +} + +type StateBase struct { + BaseHash string + Version string + Ts int64 + Data []byte +} + +type StateDiff struct { + DiffHash string + Ts int64 + BaseHash string + DiffHashArr []string + Data []byte +} + +func (sd *StateDiff) FromMap(m map[string]interface{}) bool { + quickSetStr(&sd.DiffHash, m, "diffhash") + quickSetInt64(&sd.Ts, m, "ts") + quickSetStr(&sd.BaseHash, m, "basehash") + quickSetJsonArr(&sd.DiffHashArr, m, "diffhasharr") + quickSetBytes(&sd.Data, m, "data") + return true +} + +func (sd *StateDiff) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["diffhash"] = sd.DiffHash + rtn["ts"] = sd.Ts + rtn["basehash"] = sd.BaseHash + rtn["diffhasharr"] = quickJsonArr(sd.DiffHashArr) + rtn["data"] = sd.Data + return rtn +} + +func FeStateFromShellState(state *packet.ShellState) map[string]string { + if state == nil { + return nil + } + rtn := make(map[string]string) + rtn["cwd"] = state.Cwd + envMap := shexec.EnvMapFromState(state) + if envMap["VIRTUAL_ENV"] != "" { + rtn["VIRTUAL_ENV"] = envMap["VIRTUAL_ENV"] + } + for key, val := range envMap { + if strings.HasPrefix(key, "PROMPTVAR_") { + rtn[key] = val + } + } + return rtn +} + +func (ri *RemoteInstance) FromMap(m map[string]interface{}) bool { + quickSetStr(&ri.RIId, m, "riid") + quickSetStr(&ri.Name, m, "name") + quickSetStr(&ri.SessionId, m, "sessionid") + quickSetStr(&ri.ScreenId, m, "screenid") + quickSetStr(&ri.RemoteOwnerId, m, "remoteownerid") + quickSetStr(&ri.RemoteId, m, "remoteid") + quickSetJson(&ri.FeState, m, "festate") + quickSetStr(&ri.StateBaseHash, m, "statebasehash") + quickSetJsonArr(&ri.StateDiffHashArr, m, "statediffhasharr") + return true +} + +func (ri *RemoteInstance) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["riid"] = ri.RIId + rtn["name"] = ri.Name + rtn["sessionid"] = ri.SessionId + rtn["screenid"] = ri.ScreenId + rtn["remoteownerid"] = ri.RemoteOwnerId + rtn["remoteid"] = ri.RemoteId + rtn["festate"] = quickJson(ri.FeState) + rtn["statebasehash"] = ri.StateBaseHash + rtn["statediffhasharr"] = quickJsonArr(ri.StateDiffHashArr) + return rtn +} + +type ScreenUpdateType struct { + UpdateId int64 `json:"updateid"` + ScreenId string `json:"screenid"` + LineId string `json:"lineid"` + UpdateType string `json:"updatetype"` + UpdateTs int64 `json:"updatets"` +} + +func (ScreenUpdateType) UseDBMap() {} + +type LineType struct { + ScreenId string `json:"screenid"` + UserId string `json:"userid"` + LineId string `json:"lineid"` + Ts int64 `json:"ts"` + LineNum int64 `json:"linenum"` + LineNumTemp bool `json:"linenumtemp,omitempty"` + LineLocal bool `json:"linelocal"` + LineType string `json:"linetype"` + LineState map[string]any `json:"linestate"` + Renderer string `json:"renderer,omitempty"` + Text string `json:"text,omitempty"` + Ephemeral bool `json:"ephemeral,omitempty"` + ContentHeight int64 `json:"contentheight,omitempty"` + Star bool `json:"star,omitempty"` + Archived bool `json:"archived,omitempty"` + Remove bool `json:"remove,omitempty"` +} + +func (LineType) UseDBMap() {} + +type OpenAIUsage struct { + PromptTokens int `json:"prompt_tokens"` + CompletionTokens int `json:"completion_tokens"` + TotalTokens int `json:"total_tokens"` +} + +type OpenAIChoiceType struct { + Text string `json:"text"` + Index int `json:"index"` + FinishReason string `json:"finish_reason"` +} + +type OpenAIResponse struct { + Model string `json:"model"` + Created int64 `json:"created"` + Usage *OpenAIUsage `json:"usage,omitempty"` + Choices []OpenAIChoiceType `json:"choices,omitempty"` +} + +type OpenAIPromptMessageType struct { + Role string `json:"role"` + Content string `json:"content"` + Name string `json:"name,omitempty"` +} + +type PlaybookType struct { + PlaybookId string `json:"playbookid"` + PlaybookName string `json:"playbookname"` + Description string `json:"description"` + EntryIds []string `json:"entryids"` + + // this is not persisted to DB, just for transport to FE + Entries []*PlaybookEntry `json:"entries"` +} + +func (p *PlaybookType) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["playbookid"] = p.PlaybookId + rtn["playbookname"] = p.PlaybookName + rtn["description"] = p.Description + rtn["entryids"] = quickJsonArr(p.EntryIds) + return rtn +} + +func (p *PlaybookType) FromMap(m map[string]interface{}) bool { + quickSetStr(&p.PlaybookId, m, "playbookid") + quickSetStr(&p.PlaybookName, m, "playbookname") + quickSetStr(&p.Description, m, "description") + quickSetJsonArr(&p.Entries, m, "entries") + return true +} + +// reorders p.Entries to match p.EntryIds +func (p *PlaybookType) OrderEntries() { + if len(p.Entries) == 0 { + return + } + m := make(map[string]*PlaybookEntry) + for _, entry := range p.Entries { + m[entry.EntryId] = entry + } + newList := make([]*PlaybookEntry, 0, len(p.EntryIds)) + for _, entryId := range p.EntryIds { + entry := m[entryId] + if entry != nil { + newList = append(newList, entry) + } + } + p.Entries = newList +} + +// removes from p.EntryIds (not from p.Entries) +func (p *PlaybookType) RemoveEntry(entryIdToRemove string) { + if len(p.EntryIds) == 0 { + return + } + newList := make([]string, 0, len(p.EntryIds)-1) + for _, entryId := range p.EntryIds { + if entryId == entryIdToRemove { + continue + } + newList = append(newList, entryId) + } + p.EntryIds = newList +} + +type PlaybookEntry struct { + PlaybookId string `json:"playbookid"` + EntryId string `json:"entryid"` + Alias string `json:"alias"` + CmdStr string `json:"cmdstr"` + UpdatedTs int64 `json:"updatedts"` + CreatedTs int64 `json:"createdts"` + Description string `json:"description"` + Remove bool `json:"remove,omitempty"` +} + +type BookmarkType struct { + BookmarkId string `json:"bookmarkid"` + CreatedTs int64 `json:"createdts"` + CmdStr string `json:"cmdstr"` + Alias string `json:"alias,omitempty"` + Tags []string `json:"tags"` + Description string `json:"description"` + OrderIdx int64 `json:"orderidx"` + Remove bool `json:"remove,omitempty"` +} + +func (bm *BookmarkType) GetSimpleKey() string { + return bm.BookmarkId +} + +func (bm *BookmarkType) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["bookmarkid"] = bm.BookmarkId + rtn["createdts"] = bm.CreatedTs + rtn["cmdstr"] = bm.CmdStr + rtn["alias"] = bm.Alias + rtn["description"] = bm.Description + rtn["tags"] = quickJsonArr(bm.Tags) + return rtn +} + +func (bm *BookmarkType) FromMap(m map[string]interface{}) bool { + quickSetStr(&bm.BookmarkId, m, "bookmarkid") + quickSetInt64(&bm.CreatedTs, m, "createdts") + quickSetStr(&bm.Alias, m, "alias") + quickSetStr(&bm.CmdStr, m, "cmdstr") + quickSetStr(&bm.Description, m, "description") + quickSetJsonArr(&bm.Tags, m, "tags") + return true +} + +type ResolveItem struct { + Name string + Num int + Id string + Hidden bool +} + +type SSHOpts struct { + Local bool `json:"local,omitempty"` + IsSudo bool `json:"issudo,omitempty"` + SSHHost string `json:"sshhost"` + SSHUser string `json:"sshuser"` + SSHOptsStr string `json:"sshopts,omitempty"` + SSHIdentity string `json:"sshidentity,omitempty"` + SSHPort int `json:"sshport,omitempty"` + SSHPassword string `json:"sshpassword,omitempty"` +} + +func (opts SSHOpts) GetAuthType() string { + if opts.SSHPassword != "" && opts.SSHIdentity != "" { + return RemoteAuthTypeKeyPassword + } + if opts.SSHIdentity != "" { + return RemoteAuthTypeKey + } + if opts.SSHPassword != "" { + return RemoteAuthTypePassword + } + return RemoteAuthTypeNone +} + +type RemoteOptsType struct { + Color string `json:"color"` +} + +type OpenAIOptsType struct { + Model string `json:"model"` + APIToken string `json:"apitoken"` + MaxTokens int `json:"maxtokens,omitempty"` + MaxChoices int `json:"maxchoices,omitempty"` +} + +type RemoteType struct { + RemoteId string `json:"remoteid"` + RemoteType string `json:"remotetype"` + RemoteAlias string `json:"remotealias"` + RemoteCanonicalName string `json:"remotecanonicalname"` + RemoteOpts *RemoteOptsType `json:"remoteopts"` + LastConnectTs int64 `json:"lastconnectts"` + RemoteIdx int64 `json:"remoteidx"` + Archived bool `json:"archived"` + + // SSH fields + Local bool `json:"local"` + RemoteUser string `json:"remoteuser"` + RemoteHost string `json:"remotehost"` + ConnectMode string `json:"connectmode"` + AutoInstall bool `json:"autoinstall"` + SSHOpts *SSHOpts `json:"sshopts"` + StateVars map[string]string `json:"statevars"` + + // OpenAI fields + OpenAIOpts *OpenAIOptsType `json:"openaiopts,omitempty"` +} + +func (r *RemoteType) IsSudo() bool { + return r.SSHOpts != nil && r.SSHOpts.IsSudo +} + +func (r *RemoteType) GetName() string { + if r.RemoteAlias != "" { + return r.RemoteAlias + } + return r.RemoteCanonicalName +} + +type CmdType struct { + ScreenId string `json:"screenid"` + LineId string `json:"lineid"` + Remote RemotePtrType `json:"remote"` + CmdStr string `json:"cmdstr"` + RawCmdStr string `json:"rawcmdstr"` + FeState map[string]string `json:"festate"` + StatePtr ShellStatePtr `json:"state"` + TermOpts TermOpts `json:"termopts"` + OrigTermOpts TermOpts `json:"origtermopts"` + Status string `json:"status"` + CmdPid int `json:"cmdpid"` + RemotePid int `json:"remotepid"` + DoneTs int64 `json:"donets"` + ExitCode int `json:"exitcode"` + DurationMs int `json:"durationms"` + RunOut []packet.PacketType `json:"runout,omitempty"` + RtnState bool `json:"rtnstate,omitempty"` + RtnStatePtr ShellStatePtr `json:"rtnstateptr,omitempty"` + Remove bool `json:"remove,omitempty"` +} + +func (r *RemoteType) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["remoteid"] = r.RemoteId + rtn["remotetype"] = r.RemoteType + rtn["remotealias"] = r.RemoteAlias + rtn["remotecanonicalname"] = r.RemoteCanonicalName + rtn["remoteuser"] = r.RemoteUser + rtn["remotehost"] = r.RemoteHost + rtn["connectmode"] = r.ConnectMode + rtn["autoinstall"] = r.AutoInstall + rtn["sshopts"] = quickJson(r.SSHOpts) + rtn["remoteopts"] = quickJson(r.RemoteOpts) + rtn["lastconnectts"] = r.LastConnectTs + rtn["archived"] = r.Archived + rtn["remoteidx"] = r.RemoteIdx + rtn["local"] = r.Local + rtn["statevars"] = quickJson(r.StateVars) + rtn["openaiopts"] = quickJson(r.OpenAIOpts) + return rtn +} + +func (r *RemoteType) FromMap(m map[string]interface{}) bool { + quickSetStr(&r.RemoteId, m, "remoteid") + quickSetStr(&r.RemoteType, m, "remotetype") + quickSetStr(&r.RemoteAlias, m, "remotealias") + quickSetStr(&r.RemoteCanonicalName, m, "remotecanonicalname") + quickSetStr(&r.RemoteUser, m, "remoteuser") + quickSetStr(&r.RemoteHost, m, "remotehost") + quickSetStr(&r.ConnectMode, m, "connectmode") + quickSetBool(&r.AutoInstall, m, "autoinstall") + quickSetJson(&r.SSHOpts, m, "sshopts") + quickSetJson(&r.RemoteOpts, m, "remoteopts") + quickSetInt64(&r.LastConnectTs, m, "lastconnectts") + quickSetBool(&r.Archived, m, "archived") + quickSetInt64(&r.RemoteIdx, m, "remoteidx") + quickSetBool(&r.Local, m, "local") + quickSetJson(&r.StateVars, m, "statevars") + quickSetJson(&r.OpenAIOpts, m, "openaiopts") + return true +} + +func (cmd *CmdType) ToMap() map[string]interface{} { + rtn := make(map[string]interface{}) + rtn["screenid"] = cmd.ScreenId + rtn["lineid"] = cmd.LineId + rtn["remoteownerid"] = cmd.Remote.OwnerId + rtn["remoteid"] = cmd.Remote.RemoteId + rtn["remotename"] = cmd.Remote.Name + rtn["cmdstr"] = cmd.CmdStr + rtn["rawcmdstr"] = cmd.RawCmdStr + rtn["festate"] = quickJson(cmd.FeState) + rtn["statebasehash"] = cmd.StatePtr.BaseHash + rtn["statediffhasharr"] = quickJsonArr(cmd.StatePtr.DiffHashArr) + rtn["termopts"] = quickJson(cmd.TermOpts) + rtn["origtermopts"] = quickJson(cmd.OrigTermOpts) + rtn["status"] = cmd.Status + rtn["cmdpid"] = cmd.CmdPid + rtn["remotepid"] = cmd.RemotePid + rtn["donets"] = cmd.DoneTs + rtn["exitcode"] = cmd.ExitCode + rtn["durationms"] = cmd.DurationMs + rtn["runout"] = quickJson(cmd.RunOut) + rtn["rtnstate"] = cmd.RtnState + rtn["rtnbasehash"] = cmd.RtnStatePtr.BaseHash + rtn["rtndiffhasharr"] = quickJsonArr(cmd.RtnStatePtr.DiffHashArr) + return rtn +} + +func (cmd *CmdType) FromMap(m map[string]interface{}) bool { + quickSetStr(&cmd.ScreenId, m, "screenid") + quickSetStr(&cmd.LineId, m, "lineid") + quickSetStr(&cmd.Remote.OwnerId, m, "remoteownerid") + quickSetStr(&cmd.Remote.RemoteId, m, "remoteid") + quickSetStr(&cmd.Remote.Name, m, "remotename") + quickSetStr(&cmd.CmdStr, m, "cmdstr") + quickSetStr(&cmd.RawCmdStr, m, "rawcmdstr") + quickSetJson(&cmd.FeState, m, "festate") + quickSetStr(&cmd.StatePtr.BaseHash, m, "statebasehash") + quickSetJsonArr(&cmd.StatePtr.DiffHashArr, m, "statediffhasharr") + quickSetJson(&cmd.TermOpts, m, "termopts") + quickSetJson(&cmd.OrigTermOpts, m, "origtermopts") + quickSetStr(&cmd.Status, m, "status") + quickSetInt(&cmd.CmdPid, m, "cmdpid") + quickSetInt(&cmd.RemotePid, m, "remotepid") + quickSetInt64(&cmd.DoneTs, m, "donets") + quickSetInt(&cmd.ExitCode, m, "exitcode") + quickSetInt(&cmd.DurationMs, m, "durationms") + quickSetJson(&cmd.RunOut, m, "runout") + quickSetBool(&cmd.RtnState, m, "rtnstate") + quickSetStr(&cmd.RtnStatePtr.BaseHash, m, "rtnbasehash") + quickSetJsonArr(&cmd.RtnStatePtr.DiffHashArr, m, "rtndiffhasharr") + return true +} + +func (cmd *CmdType) IsRunning() bool { + return cmd.Status == CmdStatusRunning || cmd.Status == CmdStatusDetached +} + +func makeNewLineCmd(screenId string, userId string, lineId string, renderer string, lineState map[string]any) *LineType { + rtn := &LineType{} + rtn.ScreenId = screenId + rtn.UserId = userId + rtn.LineId = lineId + rtn.Ts = time.Now().UnixMilli() + rtn.LineLocal = true + rtn.LineType = LineTypeCmd + rtn.LineId = lineId + rtn.ContentHeight = LineNoHeight + rtn.Renderer = renderer + if lineState == nil { + lineState = make(map[string]any) + } + rtn.LineState = lineState + return rtn +} + +func makeNewLineText(screenId string, userId string, text string) *LineType { + rtn := &LineType{} + rtn.ScreenId = screenId + rtn.UserId = userId + rtn.LineId = scbase.GenPromptUUID() + rtn.Ts = time.Now().UnixMilli() + rtn.LineLocal = true + rtn.LineType = LineTypeText + rtn.Text = text + rtn.ContentHeight = LineNoHeight + rtn.LineState = make(map[string]any) + return rtn +} + +func makeNewLineOpenAI(screenId string, userId string, lineId string) *LineType { + rtn := &LineType{} + rtn.ScreenId = screenId + rtn.UserId = userId + rtn.LineId = lineId + rtn.Ts = time.Now().UnixMilli() + rtn.LineLocal = true + rtn.LineType = LineTypeOpenAI + rtn.ContentHeight = LineNoHeight + rtn.Renderer = CmdRendererOpenAI + rtn.LineState = make(map[string]any) + return rtn +} + +func AddCommentLine(ctx context.Context, screenId string, userId string, commentText string) (*LineType, error) { + rtnLine := makeNewLineText(screenId, userId, commentText) + err := InsertLine(ctx, rtnLine, nil) + if err != nil { + return nil, err + } + return rtnLine, nil +} + +func AddOpenAILine(ctx context.Context, screenId string, userId string, cmd *CmdType) (*LineType, error) { + rtnLine := makeNewLineOpenAI(screenId, userId, cmd.LineId) + err := InsertLine(ctx, rtnLine, cmd) + if err != nil { + return nil, err + } + return rtnLine, nil +} + +func AddCmdLine(ctx context.Context, screenId string, userId string, cmd *CmdType, renderer string, lineState map[string]any) (*LineType, error) { + rtnLine := makeNewLineCmd(screenId, userId, cmd.LineId, renderer, lineState) + err := InsertLine(ctx, rtnLine, cmd) + if err != nil { + return nil, err + } + return rtnLine, nil +} + +func EnsureLocalRemote(ctx context.Context) error { + remote, err := GetLocalRemote(ctx) + if err != nil { + return fmt.Errorf("getting local remote from db: %w", err) + } + if remote != nil { + return nil + } + hostName, err := os.Hostname() + if err != nil { + return fmt.Errorf("getting hostname: %w", err) + } + user, err := user.Current() + if err != nil { + return fmt.Errorf("getting user: %w", err) + } + // create the local remote + localRemote := &RemoteType{ + RemoteId: scbase.GenPromptUUID(), + RemoteType: RemoteTypeSsh, + RemoteAlias: LocalRemoteAlias, + RemoteCanonicalName: fmt.Sprintf("%s@%s", user.Username, hostName), + RemoteUser: user.Username, + RemoteHost: hostName, + ConnectMode: ConnectModeStartup, + AutoInstall: true, + SSHOpts: &SSHOpts{Local: true}, + Local: true, + } + err = UpsertRemote(ctx, localRemote) + if err != nil { + return err + } + log.Printf("[db] added local remote '%s', id=%s\n", localRemote.RemoteCanonicalName, localRemote.RemoteId) + sudoRemote := &RemoteType{ + RemoteId: scbase.GenPromptUUID(), + RemoteType: RemoteTypeSsh, + RemoteAlias: "sudo", + RemoteCanonicalName: fmt.Sprintf("sudo@%s@%s", user.Username, hostName), + RemoteUser: "root", + RemoteHost: hostName, + ConnectMode: ConnectModeManual, + AutoInstall: true, + SSHOpts: &SSHOpts{Local: true, IsSudo: true}, + RemoteOpts: &RemoteOptsType{Color: "red"}, + Local: true, + } + err = UpsertRemote(ctx, sudoRemote) + if err != nil { + return err + } + log.Printf("[db] added sudo remote '%s', id=%s\n", sudoRemote.RemoteCanonicalName, sudoRemote.RemoteId) + return nil +} + +func EnsureDefaultSession(ctx context.Context) (*SessionType, error) { + session, err := GetSessionByName(ctx, DefaultSessionName) + if err != nil { + return nil, err + } + if session != nil { + return session, nil + } + _, err = InsertSessionWithName(ctx, DefaultSessionName, true) + if err != nil { + return nil, err + } + return GetSessionByName(ctx, DefaultSessionName) +} + +func createClientData(tx *TxWrap) error { + curve := elliptic.P384() + pkey, err := ecdsa.GenerateKey(curve, rand.Reader) + if err != nil { + return fmt.Errorf("generating P-834 key: %w", err) + } + pkBytes, err := x509.MarshalECPrivateKey(pkey) + if err != nil { + return fmt.Errorf("marshaling (pkcs8) private key bytes: %w", err) + } + pubBytes, err := x509.MarshalPKIXPublicKey(&pkey.PublicKey) + if err != nil { + return fmt.Errorf("marshaling (pkix) public key bytes: %w", err) + } + c := ClientData{ + ClientId: uuid.New().String(), + UserId: uuid.New().String(), + UserPrivateKeyBytes: pkBytes, + UserPublicKeyBytes: pubBytes, + ActiveSessionId: "", + WinSize: ClientWinSizeType{}, + CmdStoreType: CmdStoreTypeScreen, + } + query := `INSERT INTO client ( clientid, userid, activesessionid, userpublickeybytes, userprivatekeybytes, winsize, cmdstoretype) + VALUES (:clientid,:userid,:activesessionid,:userpublickeybytes,:userprivatekeybytes,:winsize,:cmdstoretype)` + tx.NamedExec(query, dbutil.ToDBMap(c, false)) + log.Printf("create new clientid[%s] userid[%s] with public/private keypair\n", c.ClientId, c.UserId) + return nil +} + +func EnsureClientData(ctx context.Context) (*ClientData, error) { + rtn, err := WithTxRtn(ctx, func(tx *TxWrap) (*ClientData, error) { + query := `SELECT count(*) FROM client` + count := tx.GetInt(query) + if count > 1 { + return nil, fmt.Errorf("invalid client database, multiple (%d) rows in client table", count) + } + if count == 0 { + createErr := createClientData(tx) + if createErr != nil { + return nil, createErr + } + } + cdata := dbutil.GetMappable[*ClientData](tx, `SELECT * FROM client`) + if cdata == nil { + return nil, fmt.Errorf("no client data found") + } + dbVersion := tx.GetInt(`SELECT version FROM schema_migrations`) + cdata.DBVersion = dbVersion + return cdata, nil + }) + if err != nil { + return nil, err + } + if rtn.UserId == "" { + return nil, fmt.Errorf("invalid client data (no userid)") + } + if len(rtn.UserPrivateKeyBytes) == 0 || len(rtn.UserPublicKeyBytes) == 0 { + return nil, fmt.Errorf("invalid client data (no public/private keypair)") + } + rtn.UserPrivateKey, err = x509.ParseECPrivateKey(rtn.UserPrivateKeyBytes) + if err != nil { + return nil, fmt.Errorf("invalid client data, cannot parse private key: %w", err) + } + pubKey, err := x509.ParsePKIXPublicKey(rtn.UserPublicKeyBytes) + if err != nil { + return nil, fmt.Errorf("invalid client data, cannot parse public key: %w", err) + } + var ok bool + rtn.UserPublicKey, ok = pubKey.(*ecdsa.PublicKey) + if !ok { + return nil, fmt.Errorf("invalid client data, wrong public key type: %T", pubKey) + } + return rtn, nil +} + +func SetClientOpts(ctx context.Context, clientOpts ClientOptsType) error { + txErr := WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE client SET clientopts = ?` + tx.Exec(query, quickJson(clientOpts)) + return nil + }) + return txErr +} diff --git a/wavesrv/pkg/sstore/sstore_migrate.go b/wavesrv/pkg/sstore/sstore_migrate.go new file mode 100644 index 00000000..b77bacd7 --- /dev/null +++ b/wavesrv/pkg/sstore/sstore_migrate.go @@ -0,0 +1,154 @@ +package sstore + +import ( + "context" + "fmt" + "log" + "os" + "time" + + "github.com/commandlinedev/prompt-server/pkg/scbase" +) + +const MigrationChunkSize = 10 + +type cmdMigration13Type struct { + SessionId string + ScreenId string + CmdId string +} + +type cmdMigration20Type struct { + ScreenId string + LineId string + CmdId string +} + +func getSliceChunk[T any](slice []T, chunkSize int) ([]T, []T) { + if chunkSize >= len(slice) { + return slice, nil + } + return slice[0:chunkSize], slice[chunkSize:] +} + +func RunMigration20() error { + ctx := context.Background() + startTime := time.Now() + var migrations []cmdMigration20Type + txErr := WithTx(ctx, func(tx *TxWrap) error { + tx.Select(&migrations, `SELECT * FROM cmd_migrate20`) + return nil + }) + if txErr != nil { + return fmt.Errorf("trying to get cmd20 migrations: %w", txErr) + } + log.Printf("[db] got %d cmd-line migrations\n", len(migrations)) + for len(migrations) > 0 { + var mchunk []cmdMigration20Type + mchunk, migrations = getSliceChunk(migrations, MigrationChunkSize) + err := processMigration20Chunk(ctx, mchunk) + if err != nil { + return fmt.Errorf("cmd migration failed on chunk: %w", err) + } + } + log.Printf("[db] cmd line migration done: %v\n", time.Since(startTime)) + return nil +} + +func processMigration20Chunk(ctx context.Context, mchunk []cmdMigration20Type) error { + for _, mig := range mchunk { + newFile, err := scbase.PtyOutFile(mig.ScreenId, mig.LineId) + if err != nil { + log.Printf("ptyoutfile(lineid) error: %v\n", err) + continue + } + oldFile, err := scbase.PtyOutFile(mig.ScreenId, mig.CmdId) + if err != nil { + log.Printf("ptyoutfile(cmdid) error: %v\n", err) + continue + } + err = os.Rename(oldFile, newFile) + if err != nil { + log.Printf("error renaming %s => %s: %v\n", oldFile, newFile, err) + continue + } + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + for _, mig := range mchunk { + query := `DELETE FROM cmd_migrate20 WHERE cmdid = ?` + tx.Exec(query, mig.CmdId) + } + return nil + }) + if txErr != nil { + return txErr + } + return nil +} + +func RunMigration13() error { + ctx := context.Background() + startTime := time.Now() + var migrations []cmdMigration13Type + txErr := WithTx(ctx, func(tx *TxWrap) error { + tx.Select(&migrations, `SELECT * FROM cmd_migrate`) + return nil + }) + if txErr != nil { + return fmt.Errorf("trying to get cmd13 migrations: %w", txErr) + } + log.Printf("[db] got %d cmd-screen migrations\n", len(migrations)) + for len(migrations) > 0 { + var mchunk []cmdMigration13Type + mchunk, migrations = getSliceChunk(migrations, MigrationChunkSize) + err := processMigration13Chunk(ctx, mchunk) + if err != nil { + return fmt.Errorf("cmd migration failed on chunk: %w", err) + } + } + err := os.RemoveAll(scbase.GetSessionsDir()) + if err != nil { + return fmt.Errorf("cannot remove old sessions dir %s: %w\n", scbase.GetSessionsDir(), err) + } + txErr = WithTx(ctx, func(tx *TxWrap) error { + query := `UPDATE client SET cmdstoretype = 'screen'` + tx.Exec(query) + return nil + }) + if txErr != nil { + return fmt.Errorf("cannot change client cmdstoretype: %w", err) + } + log.Printf("[db] cmd screen migration done: %v\n", time.Since(startTime)) + return nil +} + +func processMigration13Chunk(ctx context.Context, mchunk []cmdMigration13Type) error { + for _, mig := range mchunk { + newFile, err := scbase.PtyOutFile(mig.ScreenId, mig.CmdId) + if err != nil { + log.Printf("ptyoutfile error: %v\n", err) + continue + } + oldFile, err := scbase.PtyOutFile_Sessions(mig.SessionId, mig.CmdId) + if err != nil { + log.Printf("ptyoutfile_sessions error: %v\n", err) + continue + } + err = os.Rename(oldFile, newFile) + if err != nil { + log.Printf("error renaming %s => %s: %v\n", oldFile, newFile, err) + continue + } + } + txErr := WithTx(ctx, func(tx *TxWrap) error { + for _, mig := range mchunk { + query := `DELETE FROM cmd_migrate WHERE cmdid = ?` + tx.Exec(query, mig.CmdId) + } + return nil + }) + if txErr != nil { + return txErr + } + return nil +} diff --git a/wavesrv/pkg/sstore/updatebus.go b/wavesrv/pkg/sstore/updatebus.go new file mode 100644 index 00000000..72316b3a --- /dev/null +++ b/wavesrv/pkg/sstore/updatebus.go @@ -0,0 +1,228 @@ +package sstore + +import ( + "fmt" + "log" + "sync" +) + +var MainBus *UpdateBus = MakeUpdateBus() + +const PtyDataUpdateStr = "pty" +const ModelUpdateStr = "model" +const UpdateChSize = 100 + +type UpdatePacket interface { + UpdateType() string + Clean() +} + +type PtyDataUpdate struct { + ScreenId string `json:"screenid,omitempty"` + LineId string `json:"lineid,omitempty"` + RemoteId string `json:"remoteid,omitempty"` + PtyPos int64 `json:"ptypos"` + PtyData64 string `json:"ptydata64"` + PtyDataLen int64 `json:"ptydatalen"` +} + +func (*PtyDataUpdate) UpdateType() string { + return PtyDataUpdateStr +} + +func (pdu *PtyDataUpdate) Clean() {} + +type ModelUpdate struct { + Sessions []*SessionType `json:"sessions,omitempty"` + ActiveSessionId string `json:"activesessionid,omitempty"` + Screens []*ScreenType `json:"screens,omitempty"` + ScreenLines *ScreenLinesType `json:"screenlines,omitempty"` + Line *LineType `json:"line,omitempty"` + Lines []*LineType `json:"lines,omitempty"` + Cmd *CmdType `json:"cmd,omitempty"` + CmdLine *CmdLineType `json:"cmdline,omitempty"` + Info *InfoMsgType `json:"info,omitempty"` + ClearInfo bool `json:"clearinfo,omitempty"` + Remotes []interface{} `json:"remotes,omitempty"` // []*remote.RemoteState + History *HistoryInfoType `json:"history,omitempty"` + Interactive bool `json:"interactive"` + Connect bool `json:"connect,omitempty"` + MainView string `json:"mainview,omitempty"` + Bookmarks []*BookmarkType `json:"bookmarks,omitempty"` + SelectedBookmark string `json:"selectedbookmark,omitempty"` + HistoryViewData *HistoryViewData `json:"historyviewdata,omitempty"` + ClientData *ClientData `json:"clientdata,omitempty"` + RemoteView *RemoteViewType `json:"remoteview,omitempty"` +} + +func (*ModelUpdate) UpdateType() string { + return ModelUpdateStr +} + +func (update *ModelUpdate) Clean() { + if update == nil { + return + } + update.ClientData = update.ClientData.Clean() +} + +type RemoteViewType struct { + RemoteShowAll bool `json:"remoteshowall,omitempty"` + PtyRemoteId string `json:"ptyremoteid,omitempty"` + RemoteEdit *RemoteEditType `json:"remoteedit,omitempty"` +} + +func InfoMsgUpdate(infoMsgFmt string, args ...interface{}) *ModelUpdate { + msg := fmt.Sprintf(infoMsgFmt, args...) + return &ModelUpdate{ + Info: &InfoMsgType{InfoMsg: msg}, + } +} + +type HistoryViewData struct { + Items []*HistoryItemType `json:"items"` + Offset int `json:"offset"` + RawOffset int `json:"rawoffset"` + NextRawOffset int `json:"nextrawoffset"` + HasMore bool `json:"hasmore"` + Lines []*LineType `json:"lines"` + Cmds []*CmdType `json:"cmds"` +} + +type RemoteEditType struct { + RemoteEdit bool `json:"remoteedit"` + RemoteId string `json:"remoteid,omitempty"` + ErrorStr string `json:"errorstr,omitempty"` + InfoStr string `json:"infostr,omitempty"` + KeyStr string `json:"keystr,omitempty"` + HasPassword bool `json:"haspassword,omitempty"` +} + +type InfoMsgType struct { + InfoTitle string `json:"infotitle"` + InfoError string `json:"infoerror,omitempty"` + InfoMsg string `json:"infomsg,omitempty"` + InfoMsgHtml bool `json:"infomsghtml,omitempty"` + WebShareLink bool `json:"websharelink,omitempty"` + InfoComps []string `json:"infocomps,omitempty"` + InfoCompsMore bool `json:"infocompssmore,omitempty"` + InfoLines []string `json:"infolines,omitempty"` + TimeoutMs int64 `json:"timeoutms,omitempty"` +} + +type HistoryInfoType struct { + HistoryType string `json:"historytype"` + SessionId string `json:"sessionid,omitempty"` + ScreenId string `json:"screenid,omitempty"` + Items []*HistoryItemType `json:"items"` + Show bool `json:"show"` +} + +type CmdLineType struct { + CmdLine string `json:"cmdline"` + CursorPos int `json:"cursorpos"` +} + +type UpdateChannel struct { + ScreenId string + ClientId string + Ch chan interface{} +} + +func (uch UpdateChannel) Match(screenId string) bool { + if screenId == "" { + return true + } + return screenId == uch.ScreenId +} + +type UpdateBus struct { + Lock *sync.Mutex + Channels map[string]UpdateChannel +} + +func MakeUpdateBus() *UpdateBus { + return &UpdateBus{ + Lock: &sync.Mutex{}, + Channels: make(map[string]UpdateChannel), + } +} + +// always returns a new channel +func (bus *UpdateBus) RegisterChannel(clientId string, screenId string) chan interface{} { + bus.Lock.Lock() + defer bus.Lock.Unlock() + uch, found := bus.Channels[clientId] + if found { + close(uch.Ch) + uch.ScreenId = screenId + uch.Ch = make(chan interface{}, UpdateChSize) + } else { + uch = UpdateChannel{ + ClientId: clientId, + ScreenId: screenId, + Ch: make(chan interface{}, UpdateChSize), + } + } + bus.Channels[clientId] = uch + return uch.Ch +} + +func (bus *UpdateBus) UnregisterChannel(clientId string) { + bus.Lock.Lock() + defer bus.Lock.Unlock() + uch, found := bus.Channels[clientId] + if found { + close(uch.Ch) + delete(bus.Channels, clientId) + } +} + +func (bus *UpdateBus) SendUpdate(update UpdatePacket) { + if update == nil { + return + } + update.Clean() + bus.Lock.Lock() + defer bus.Lock.Unlock() + for _, uch := range bus.Channels { + select { + case uch.Ch <- update: + + default: + log.Printf("[error] dropped update on updatebus uch clientid=%s\n", uch.ClientId) + } + } +} + +func (bus *UpdateBus) SendScreenUpdate(screenId string, update UpdatePacket) { + if update == nil { + return + } + update.Clean() + bus.Lock.Lock() + defer bus.Lock.Unlock() + for _, uch := range bus.Channels { + if uch.Match(screenId) { + select { + case uch.Ch <- update: + + default: + log.Printf("[error] dropped update on updatebus uch clientid=%s\n", uch.ClientId) + } + } + } +} + +func MakeSessionsUpdateForRemote(sessionId string, ri *RemoteInstance) []*SessionType { + return []*SessionType{ + &SessionType{ + SessionId: sessionId, + Remotes: []*RemoteInstance{ri}, + }, + } +} + +type BookmarksViewType struct { + Bookmarks []*BookmarkType `json:"bookmarks"` +} diff --git a/wavesrv/pkg/utilfn/linediff.go b/wavesrv/pkg/utilfn/linediff.go new file mode 100644 index 00000000..5d4a2e9e --- /dev/null +++ b/wavesrv/pkg/utilfn/linediff.go @@ -0,0 +1,132 @@ +package utilfn + +import ( + "bytes" + "encoding/binary" + "fmt" + "strings" +) + +const LineDiffVersion = 0 + +type LineDiffType struct { + Lines []int + NewData []string +} + +// simple encoding +// a 0 means read a line from NewData +// a non-zero number means read the 1-indexed line from OldData +func applyDiff(oldData []string, diff LineDiffType) ([]string, error) { + rtn := make([]string, 0, len(diff.Lines)) + newDataPos := 0 + for i := 0; i < len(diff.Lines); i++ { + if diff.Lines[i] == 0 { + if newDataPos >= len(diff.NewData) { + return nil, fmt.Errorf("not enough newdata for diff") + } + rtn = append(rtn, diff.NewData[newDataPos]) + newDataPos++ + } else { + idx := diff.Lines[i] - 1 // 1-indexed + if idx < 0 || idx >= len(oldData) { + return nil, fmt.Errorf("diff index out of bounds %d old-data-len:%d", idx, len(oldData)) + } + rtn = append(rtn, oldData[idx]) + } + } + return rtn, nil +} + +func putUVarint(buf *bytes.Buffer, viBuf []byte, ival int) { + l := binary.PutUvarint(viBuf, uint64(ival)) + buf.Write(viBuf[0:l]) +} + +// simple encoding +// write varints. first version, then len, then len-number-of-varints, then fill the rest with newdata +// [version] [len-varint] [varint]xlen... newdata (bytes) +func encodeDiff(diff LineDiffType) []byte { + var buf bytes.Buffer + viBuf := make([]byte, binary.MaxVarintLen64) + putUVarint(&buf, viBuf, 0) + putUVarint(&buf, viBuf, len(diff.Lines)) + for _, val := range diff.Lines { + putUVarint(&buf, viBuf, val) + } + for _, str := range diff.NewData { + buf.WriteString(str) + buf.WriteByte('\n') + } + return buf.Bytes() +} + +func decodeDiff(diffBytes []byte) (LineDiffType, error) { + var rtn LineDiffType + r := bytes.NewBuffer(diffBytes) + version, err := binary.ReadUvarint(r) + if err != nil { + return rtn, fmt.Errorf("invalid diff, cannot read version: %v", err) + } + if version != LineDiffVersion { + return rtn, fmt.Errorf("invalid diff, bad version: %d", version) + } + linesLen64, err := binary.ReadUvarint(r) + if err != nil { + return rtn, fmt.Errorf("invalid diff, cannot read lines length: %v", err) + } + linesLen := int(linesLen64) + rtn.Lines = make([]int, linesLen) + for idx := 0; idx < linesLen; idx++ { + vi, err := binary.ReadUvarint(r) + if err != nil { + return rtn, fmt.Errorf("invalid diff, cannot read line %d: %v", idx, err) + } + rtn.Lines[idx] = int(vi) + } + restOfInput := string(r.Bytes()) + rtn.NewData = strings.Split(restOfInput, "\n") + return rtn, nil +} + +func makeDiff(oldData []string, newData []string) LineDiffType { + var rtn LineDiffType + oldDataMap := make(map[string]int) // 1-indexed + for idx, str := range oldData { + if _, found := oldDataMap[str]; found { + continue + } + oldDataMap[str] = idx + 1 + } + rtn.Lines = make([]int, len(newData)) + for idx, str := range newData { + oldIdx, found := oldDataMap[str] + if found { + rtn.Lines[idx] = oldIdx + } else { + rtn.Lines[idx] = 0 + rtn.NewData = append(rtn.NewData, str) + } + } + return rtn +} + +func MakeDiff(str1 string, str2 string) []byte { + str1Arr := strings.Split(str1, "\n") + str2Arr := strings.Split(str2, "\n") + diff := makeDiff(str1Arr, str2Arr) + return encodeDiff(diff) +} + +func ApplyDiff(str1 string, diffBytes []byte) (string, error) { + diff, err := decodeDiff(diffBytes) + if err != nil { + return "", err + } + str1Arr := strings.Split(str1, "\n") + str2Arr, err := applyDiff(str1Arr, diff) + if err != nil { + return "", err + } + return strings.Join(str2Arr, "\n"), nil +} diff --git a/wavesrv/pkg/utilfn/utilfn.go b/wavesrv/pkg/utilfn/utilfn.go new file mode 100644 index 00000000..5b3e70b1 --- /dev/null +++ b/wavesrv/pkg/utilfn/utilfn.go @@ -0,0 +1,208 @@ +package utilfn + +import ( + "crypto/sha1" + "encoding/base64" + "regexp" + "strings" + "unicode/utf8" +) + +var HexDigits = []byte{'0', '1', '2', '3', '4', '5', '6', '7', '8', '9', 'a', 'b', 'c', 'd', 'e', 'f'} + +func GetStrArr(v interface{}, field string) []string { + if v == nil { + return nil + } + m, ok := v.(map[string]interface{}) + if !ok { + return nil + } + fieldVal := m[field] + if fieldVal == nil { + return nil + } + iarr, ok := fieldVal.([]interface{}) + if !ok { + return nil + } + var sarr []string + for _, iv := range iarr { + if sv, ok := iv.(string); ok { + sarr = append(sarr, sv) + } + } + return sarr +} + +func GetBool(v interface{}, field string) bool { + if v == nil { + return false + } + m, ok := v.(map[string]interface{}) + if !ok { + return false + } + fieldVal := m[field] + if fieldVal == nil { + return false + } + bval, ok := fieldVal.(bool) + if !ok { + return false + } + return bval +} + +var needsQuoteRe = regexp.MustCompile(`[^\w@%:,./=+-]`) + +// minimum maxlen=6 +func ShellQuote(val string, forceQuote bool, maxLen int) string { + if maxLen < 6 { + maxLen = 6 + } + rtn := val + if needsQuoteRe.MatchString(val) { + rtn = "'" + strings.ReplaceAll(val, "'", `'"'"'`) + "'" + } + if strings.HasPrefix(rtn, "\"") || strings.HasPrefix(rtn, "'") { + if len(rtn) > maxLen { + return rtn[0:maxLen-4] + "..." + rtn[0:1] + } + return rtn + } + if forceQuote { + if len(rtn) > maxLen-2 { + return "\"" + rtn[0:maxLen-5] + "...\"" + } + return "\"" + rtn + "\"" + } else { + if len(rtn) > maxLen { + return rtn[0:maxLen-3] + "..." + } + return rtn + } +} + +func EllipsisStr(s string, maxLen int) string { + if maxLen < 4 { + maxLen = 4 + } + if len(s) > maxLen { + return s[0:maxLen-3] + "..." + } + return s +} + +func LongestPrefix(root string, strs []string) string { + if len(strs) == 0 { + return root + } + if len(strs) == 1 { + comp := strs[0] + if len(comp) >= len(root) && strings.HasPrefix(comp, root) { + if strings.HasSuffix(comp, "/") { + return strs[0] + } + return strs[0] + } + } + lcp := strs[0] + for i := 1; i < len(strs); i++ { + s := strs[i] + for j := 0; j < len(lcp); j++ { + if j >= len(s) || lcp[j] != s[j] { + lcp = lcp[0:j] + break + } + } + } + if len(lcp) < len(root) || !strings.HasPrefix(lcp, root) { + return root + } + return lcp +} + +func ContainsStr(strs []string, test string) bool { + for _, s := range strs { + if s == test { + return true + } + } + return false +} + +func IsPrefix(strs []string, test string) bool { + for _, s := range strs { + if len(s) > len(test) && strings.HasPrefix(s, test) { + return true + } + } + return false +} + +type StrWithPos struct { + Str string + Pos int // this is a 'rune' position (not a byte position) +} + +func (sp StrWithPos) String() string { + return strWithCursor(sp.Str, sp.Pos) +} + +func ParseToSP(s string) StrWithPos { + idx := strings.Index(s, "[*]") + if idx == -1 { + return StrWithPos{Str: s} + } + return StrWithPos{Str: s[0:idx] + s[idx+3:], Pos: utf8.RuneCountInString(s[0:idx])} +} + +func strWithCursor(str string, pos int) string { + if pos < 0 { + return "[*]_" + str + } + if pos >= len(str) { + if pos > len(str) { + return str + "_[*]" + } + return str + "[*]" + } + + var rtn []rune + for _, ch := range str { + if len(rtn) == pos { + rtn = append(rtn, '[', '*', ']') + } + rtn = append(rtn, ch) + } + return string(rtn) +} + +func (sp StrWithPos) Prepend(str string) StrWithPos { + return StrWithPos{Str: str + sp.Str, Pos: utf8.RuneCountInString(str) + sp.Pos} +} + +func (sp StrWithPos) Append(str string) StrWithPos { + return StrWithPos{Str: sp.Str + str, Pos: sp.Pos} +} + +// returns base64 hash of data +func Sha1Hash(data []byte) string { + hvalRaw := sha1.Sum(data) + hval := base64.StdEncoding.EncodeToString(hvalRaw[:]) + return hval +} + +func ChunkSlice[T any](s []T, chunkSize int) [][]T { + var rtn [][]T + for len(rtn) > 0 { + if len(s) <= chunkSize { + rtn = append(rtn, s) + break + } + rtn = append(rtn, s[:chunkSize]) + s = s[chunkSize:] + } + return rtn +} diff --git a/wavesrv/pkg/utilfn/utilfn_test.go b/wavesrv/pkg/utilfn/utilfn_test.go new file mode 100644 index 00000000..d5d7c81b --- /dev/null +++ b/wavesrv/pkg/utilfn/utilfn_test.go @@ -0,0 +1,48 @@ +package utilfn + +import ( + "fmt" + "testing" +) + +const Str1 = ` +hello +line #2 +more +stuff +apple +` + +const Str2 = ` +line #2 +apple +grapes +banana +` + +const Str3 = ` +more +stuff +banana +coconut +` + +func testDiff(t *testing.T, str1 string, str2 string) { + diffBytes := MakeDiff(str1, str2) + fmt.Printf("diff-len: %d\n", len(diffBytes)) + out, err := ApplyDiff(str1, diffBytes) + if err != nil { + t.Errorf("error in diff: %v", err) + return + } + if out != str2 { + t.Errorf("bad diff output") + } +} + +func TestDiff(t *testing.T) { + testDiff(t, Str1, Str2) + testDiff(t, Str2, Str3) + testDiff(t, Str1, Str3) + testDiff(t, Str3, Str1) +} diff --git a/wavesrv/pkg/wsshell/wsshell.go b/wavesrv/pkg/wsshell/wsshell.go new file mode 100644 index 00000000..d6fe69e7 --- /dev/null +++ b/wavesrv/pkg/wsshell/wsshell.go @@ -0,0 +1,189 @@ +package wsshell + +import ( + "encoding/json" + "fmt" + "log" + "net/http" + "net/url" + "sync" + "time" + + "github.com/google/uuid" + "github.com/gorilla/websocket" +) + +const readWaitTimeout = 15 * time.Second +const writeWaitTimeout = 10 * time.Second +const pingPeriodTickTime = 10 * time.Second +const initialPingTime = 1 * time.Second + +var upgrader = websocket.Upgrader{ + ReadBufferSize: 4 * 1024, + WriteBufferSize: 32 * 1024, + HandshakeTimeout: 1 * time.Second, + CheckOrigin: func(r *http.Request) bool { return true }, +} + +type WSShell struct { + Conn *websocket.Conn + RemoteAddr string + ConnId string + Query url.Values + OpenTime time.Time + NumPings int + LastPing time.Time + LastRecv time.Time + Header http.Header + + CloseChan chan bool + WriteChan chan []byte + ReadChan chan []byte +} + +func (ws *WSShell) NonBlockingWrite(data []byte) bool { + select { + case ws.WriteChan <- data: + return true + + default: + return false + } +} + +func (ws *WSShell) WritePing() error { + now := time.Now() + pingMessage := map[string]interface{}{"type": "ping", "stime": now.Unix()} + jsonVal, _ := json.Marshal(pingMessage) + _ = ws.Conn.SetWriteDeadline(time.Now().Add(writeWaitTimeout)) // no error + err := ws.Conn.WriteMessage(websocket.TextMessage, jsonVal) + ws.NumPings++ + ws.LastPing = now + if err != nil { + return err + } + return nil +} + +func (ws *WSShell) WriteJson(val interface{}) error { + if ws.IsClosed() { + return fmt.Errorf("cannot write packet, empty or closed wsshell") + } + barr, err := json.Marshal(val) + if err != nil { + return err + } + ws.WriteChan <- barr + return nil +} + +func (ws *WSShell) WritePump() { + ticker := time.NewTicker(initialPingTime) + defer func() { + ticker.Stop() + ws.Conn.Close() + }() + initialPing := true + for { + select { + case <-ticker.C: + err := ws.WritePing() + if err != nil { + log.Printf("WritePump %s err: %v\n", ws.RemoteAddr, err) + return + } + if initialPing { + initialPing = false + ticker.Reset(pingPeriodTickTime) + } + + case msgBytes, ok := <-ws.WriteChan: + if !ok { + return + } + _ = ws.Conn.SetWriteDeadline(time.Now().Add(writeWaitTimeout)) // no error + err := ws.Conn.WriteMessage(websocket.TextMessage, msgBytes) + if err != nil { + log.Printf("WritePump %s err: %v\n", ws.RemoteAddr, err) + return + } + } + } +} + +func (ws *WSShell) ReadPump() { + readWait := readWaitTimeout + defer func() { + ws.Conn.Close() + }() + ws.Conn.SetReadLimit(4096) + ws.Conn.SetReadDeadline(time.Now().Add(readWait)) + for { + _, message, err := ws.Conn.ReadMessage() + if err != nil { + log.Printf("ReadPump %s Err: %v\n", ws.RemoteAddr, err) + break + } + jmsg := map[string]interface{}{} + err = json.Unmarshal(message, &jmsg) + if err != nil { + log.Printf("Error unmarshalling json: %v\n", err) + break + } + ws.Conn.SetReadDeadline(time.Now().Add(readWait)) + ws.LastRecv = time.Now() + if str, ok := jmsg["type"].(string); ok && str == "pong" { + // nothing + continue + } + if str, ok := jmsg["type"].(string); ok && str == "ping" { + now := time.Now() + pongMessage := map[string]interface{}{"type": "pong", "stime": now.Unix()} + jsonVal, _ := json.Marshal(pongMessage) + ws.WriteChan <- jsonVal + continue + } + ws.ReadChan <- message + } +} + +func (ws *WSShell) IsClosed() bool { + select { + case <-ws.CloseChan: + return true + + default: + return false + } +} + +func StartWS(w http.ResponseWriter, r *http.Request) (*WSShell, error) { + conn, err := upgrader.Upgrade(w, r, nil) + if err != nil { + return nil, err + } + ws := WSShell{Conn: conn, ConnId: uuid.New().String(), OpenTime: time.Now()} + ws.CloseChan = make(chan bool) + ws.WriteChan = make(chan []byte, 10) + ws.ReadChan = make(chan []byte, 10) + ws.RemoteAddr = r.RemoteAddr + ws.Query = r.URL.Query() + ws.Header = r.Header + var wg sync.WaitGroup + wg.Add(1) + go func() { + defer wg.Done() + ws.WritePump() + }() + wg.Add(1) + go func() { + defer wg.Done() + ws.ReadPump() + }() + go func() { + wg.Wait() + close(ws.CloseChan) + close(ws.ReadChan) + }() + return &ws, nil +} diff --git a/wavesrv/scripthaus.md b/wavesrv/scripthaus.md new file mode 100644 index 00000000..2f9ec41c --- /dev/null +++ b/wavesrv/scripthaus.md @@ -0,0 +1,16 @@ +# SH2 Server Commands + +```bash +# @scripthaus command dump-schema-dev +sqlite3 /Users/mike/prompt-dev/prompt.db .schema > db/schema.sql +``` + +```bash +# @scripthaus command opendb-dev +sqlite3 /Users/mike/prompt-dev/prompt.db +``` + +```bash +# @scripthaus command build +go build -ldflags "-X main.BuildTime=$(date +'%Y%m%d%H%M')" -o bin/local-server ./cmd +```