This commit is contained in:
+106
-2
@@ -9,9 +9,13 @@ import (
|
||||
"log"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"kb-editor/internal/aifallback"
|
||||
"kb-editor/internal/staging"
|
||||
"kb-editor/internal/store"
|
||||
)
|
||||
|
||||
@@ -35,6 +39,16 @@ func main() {
|
||||
log.Fatalf("initialize store: %v", err)
|
||||
}
|
||||
|
||||
aiService, aiTimeout, err := aiServiceFromEnv(cfg.Mode, s.DataDir())
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
}
|
||||
if aiService != nil {
|
||||
cfg.AIFallbackEnabled = true
|
||||
cfg.AIFallbackTimeoutSeconds = int(aiTimeout.Seconds())
|
||||
cfg.AIFallbackModel = aiService.Model()
|
||||
}
|
||||
|
||||
reloadInterval, err := autoReloadInterval(cfg.Mode)
|
||||
if err != nil {
|
||||
log.Fatal(err)
|
||||
@@ -48,15 +62,19 @@ func main() {
|
||||
log.Fatal(err)
|
||||
}
|
||||
|
||||
app := newApp(s, sub, cfg)
|
||||
app := newApp(s, sub, cfg).withAI(aiService)
|
||||
handler := requestLogger(optionalBasicAuth(app.routes()))
|
||||
|
||||
writeTimeout := 60 * time.Second
|
||||
if aiService != nil && aiTimeout+30*time.Second > writeTimeout {
|
||||
writeTimeout = aiTimeout + 30*time.Second
|
||||
}
|
||||
srv := &http.Server{
|
||||
Addr: listen,
|
||||
Handler: handler,
|
||||
ReadHeaderTimeout: 10 * time.Second,
|
||||
ReadTimeout: 30 * time.Second,
|
||||
WriteTimeout: 60 * time.Second,
|
||||
WriteTimeout: writeTimeout,
|
||||
IdleTimeout: 90 * time.Second,
|
||||
}
|
||||
|
||||
@@ -66,6 +84,9 @@ func main() {
|
||||
if reloadInterval > 0 {
|
||||
log.Printf("Automatic index reload: %s", reloadInterval)
|
||||
}
|
||||
if aiService != nil {
|
||||
log.Printf("AI fallback enabled: model=%q timeout=%s staging=%s", aiService.Model(), aiTimeout, aiService.StagingDir())
|
||||
}
|
||||
if u := os.Getenv("BASIC_AUTH_USER"); u != "" {
|
||||
log.Printf("Basic authentication enabled for user %q", u)
|
||||
}
|
||||
@@ -127,6 +148,89 @@ func startAutoReload(s *store.Store, interval time.Duration) {
|
||||
}
|
||||
}
|
||||
|
||||
func aiServiceFromEnv(mode, dataDir string) (*aifallback.Service, time.Duration, error) {
|
||||
enabled, err := envBool("AI_FALLBACK_ENABLED", false)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if !enabled {
|
||||
return nil, 0, nil
|
||||
}
|
||||
if mode != "google" {
|
||||
return nil, 0, fmt.Errorf("AI_FALLBACK_ENABLED is only supported with APP_MODE=google")
|
||||
}
|
||||
|
||||
timeout, err := time.ParseDuration(envOr("OLLAMA_TIMEOUT", "10m"))
|
||||
if err != nil || timeout < time.Second {
|
||||
return nil, 0, fmt.Errorf("invalid OLLAMA_TIMEOUT: expected a duration such as 10m")
|
||||
}
|
||||
maxConcurrent, err := strconv.Atoi(envOr("OLLAMA_MAX_CONCURRENT", "1"))
|
||||
if err != nil || maxConcurrent < 1 || maxConcurrent > 16 {
|
||||
return nil, 0, fmt.Errorf("OLLAMA_MAX_CONCURRENT must be an integer between 1 and 16")
|
||||
}
|
||||
autoReply, err := envBool("OLLAMA_STAGING_AUTO_REPLY", false)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
minScore, err := strconv.ParseFloat(envOr("OLLAMA_STAGING_MIN_SCORE", "0.78"), 64)
|
||||
if err != nil || minScore < 0 || minScore > 1 {
|
||||
return nil, 0, fmt.Errorf("OLLAMA_STAGING_MIN_SCORE must be between 0 and 1")
|
||||
}
|
||||
|
||||
stagingDir := strings.TrimSpace(os.Getenv("STAGING_DIR"))
|
||||
if stagingDir == "" {
|
||||
stagingDir = filepath.Join(filepath.Dir(dataDir), "staging")
|
||||
}
|
||||
stagingAbs, err := filepath.Abs(stagingDir)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
dataAbs, err := filepath.Abs(dataDir)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
if pathContains(dataAbs, stagingAbs) || pathContains(stagingAbs, dataAbs) {
|
||||
return nil, 0, fmt.Errorf("STAGING_DIR (%s) must be separate from DATA_DIR (%s)", stagingAbs, dataAbs)
|
||||
}
|
||||
|
||||
st, err := staging.New(stagingAbs)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
svc, err := aifallback.New(aifallback.Config{
|
||||
BaseURL: envOr("OLLAMA_BASE_URL", "http://ollama:11434"),
|
||||
Model: strings.TrimSpace(os.Getenv("OLLAMA_MODEL")),
|
||||
Timeout: timeout,
|
||||
MaxConcurrent: maxConcurrent,
|
||||
AutoReply: autoReply,
|
||||
MinScore: minScore,
|
||||
}, st)
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return svc, timeout, nil
|
||||
}
|
||||
|
||||
func pathContains(parent, child string) bool {
|
||||
rel, err := filepath.Rel(parent, child)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return rel == "." || (rel != ".." && !strings.HasPrefix(rel, ".."+string(filepath.Separator)))
|
||||
}
|
||||
|
||||
func envBool(key string, fallback bool) (bool, error) {
|
||||
raw := strings.TrimSpace(os.Getenv(key))
|
||||
if raw == "" {
|
||||
return fallback, nil
|
||||
}
|
||||
value, err := strconv.ParseBool(raw)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("invalid %s %q: expected true or false", key, raw)
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
|
||||
func envOr(key, fallback string) string {
|
||||
if v := strings.TrimSpace(os.Getenv(key)); v != "" {
|
||||
return v
|
||||
|
||||
Reference in New Issue
Block a user