x

Programs, configuration and documentation that don't fit anywhere else
Log | Files | Refs | README | LICENSE

llm.go (2295B)


      1 package main
      2 
      3 import (
      4 	"bufio"
      5 	"bytes"
      6 	"errors"
      7 	"flag"
      8 	"fmt"
      9 	"io"
     10 	"log"
     11 	"net/http"
     12 	"os"
     13 	"path"
     14 
     15 	"olowe.co/x/openai"
     16 )
     17 
     18 var model = flag.String("m", "", "model")
     19 var sysPrompt = flag.String("s", "", "system prompt")
     20 var converse = flag.Bool("c", false, "start a back-and-forth chat")
     21 
     22 func copyAll(w io.Writer, paths []string) (n int64, err error) {
     23 	if len(paths) == 0 {
     24 		return io.Copy(w, os.Stdin)
     25 	}
     26 	var errs []error
     27 	for _, name := range paths {
     28 		f, err := os.Open(name)
     29 		if err != nil {
     30 			return n, err
     31 		}
     32 		nn, err := io.Copy(w, f)
     33 		if err != nil {
     34 			errs = append(errs, fmt.Errorf("copy %s: %w", name, err))
     35 		}
     36 		f.Close()
     37 		n += nn
     38 	}
     39 	return n, errors.Join(errs...)
     40 }
     41 
     42 func init() {
     43 	log.SetFlags(0)
     44 	log.SetPrefix("llm: ")
     45 	flag.Parse()
     46 }
     47 
     48 func main() {
     49 	confDir, err := os.UserConfigDir()
     50 	if err != nil {
     51 		log.Fatal(err)
     52 	}
     53 	config, err := readConfig(path.Join(confDir, "openai"))
     54 	if err != nil {
     55 		log.Fatalf("read configuration: %v", err)
     56 	}
     57 	if *model == "" {
     58 		*model = config.DefaultModel
     59 	}
     60 	if config.BaseURL == "" {
     61 		config.BaseURL = "http://127.0.0.1:8080"
     62 	}
     63 	client := &openai.Client{http.DefaultClient, config.Token, config.BaseURL}
     64 
     65 	chat := openai.Chat{Model: *model}
     66 	if *sysPrompt != "" {
     67 		chat.Messages = []openai.Message{
     68 			{openai.RoleSystem, *sysPrompt},
     69 		}
     70 	}
     71 	buf := &bytes.Buffer{}
     72 	if !*converse {
     73 		_, err := copyAll(buf, flag.Args())
     74 		if err != nil {
     75 			log.Fatalln("construct prompt:", err)
     76 		}
     77 		msg := openai.Message{openai.RoleUser, buf.String()}
     78 		chat.Messages = append(chat.Messages, msg)
     79 		reply, err := client.Complete(&chat)
     80 		if err != nil {
     81 			log.Fatalln("llm complete:", err)
     82 		}
     83 		fmt.Println(reply.Content)
     84 		return
     85 	}
     86 
     87 	sc := bufio.NewScanner(os.Stdin)
     88 	if len(flag.Args()) > 0 {
     89 		log.Println("conversation mode, ignoring arguments")
     90 	}
     91 	for sc.Scan() {
     92 		if sc.Text() == "." {
     93 			msg := openai.Message{openai.RoleUser, buf.String()}
     94 			chat.Messages = append(chat.Messages, msg)
     95 			reply, err := client.Complete(&chat)
     96 			if err != nil {
     97 				fmt.Fprintln(os.Stderr, "chat not completed:", err)
     98 				continue // try again, allowing a retry with another "." line
     99 			}
    100 			buf.Reset()
    101 			fmt.Println(reply.Content)
    102 			chat.Messages = append(chat.Messages, *reply)
    103 			continue
    104 		}
    105 		fmt.Fprintln(buf, sc.Text())
    106 	}
    107 }