rony-chat-bot/internal/config/config.go

141 lines
3.3 KiB
Go
Raw Normal View History

package config
import (
"fmt"
"os"
"strings"
"gopkg.in/yaml.v3"
)
type Server struct {
Host string `yaml:"host"`
Port int `yaml:"port"`
ReadTimeoutMS int `yaml:"read_timeout_ms"`
CORSOrigins []string `yaml:"cors_origins"`
RateLimit RateLimit `yaml:"rate_limit"`
}
type RateLimit struct {
RequestsPerMinute int `yaml:"requests_per_minute"`
Burst int `yaml:"burst"`
}
type Provider struct {
Name string `yaml:"name"`
Type string `yaml:"type"`
Model string `yaml:"model,omitempty"`
ModelPath string `yaml:"model_path,omitempty"`
Endpoint string `yaml:"endpoint,omitempty"`
ContextSize int `yaml:"context_size,omitempty"`
MaxTokens int `yaml:"max_tokens,omitempty"`
NGPULayers int `yaml:"n_gpu_layers,omitempty"`
Temperature float32 `yaml:"temperature,omitempty"`
APIKeyEnv string `yaml:"api_key_env,omitempty"`
Default bool `yaml:"default,omitempty"`
}
type RAG struct {
Enabled bool `yaml:"enabled"`
DataPath string `yaml:"data_path"`
ChunkSize int `yaml:"chunk_size"`
ChunkOverlap int `yaml:"chunk_overlap"`
DBPath string `yaml:"db_path"`
TopK int `yaml:"top_k"`
Tokenize string `yaml:"tokenize"`
}
type Persona struct {
Name string `yaml:"name"`
Tone string `yaml:"tone"`
Language string `yaml:"language"`
Constraints []string `yaml:"constraints"`
Intro string `yaml:"intro"`
}
type Logging struct {
Level string `yaml:"level"`
Format string `yaml:"format"`
Output string `yaml:"output"`
}
type Config struct {
Server Server `yaml:"server"`
Providers []Provider `yaml:"providers"`
RAG RAG `yaml:"rag"`
Persona Persona `yaml:"persona"`
SystemPrompt string `yaml:"system_prompt"`
Logging Logging `yaml:"logging"`
}
func Load(path string) (*Config, error) {
raw, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read config %s: %w", path, err)
}
expanded := os.ExpandEnv(string(raw))
var cfg Config
if err := yaml.Unmarshal([]byte(expanded), &cfg); err != nil {
return nil, fmt.Errorf("parse config %s: %w", path, err)
}
if err := cfg.validate(); err != nil {
return nil, err
}
return &cfg, nil
}
func (c *Config) validate() error {
if len(c.Providers) == 0 {
return fmt.Errorf("config: at least one provider must be configured")
}
defaultCount := 0
for _, p := range c.Providers {
if p.Default {
defaultCount++
}
}
if defaultCount == 0 {
c.Providers[0].Default = true
} else if defaultCount > 1 {
return fmt.Errorf("config: multiple providers marked as default")
}
if c.Server.Port == 0 {
c.Server.Port = 7331
}
if c.Server.ReadTimeoutMS == 0 {
c.Server.ReadTimeoutMS = 30000
}
if c.RAG.ChunkSize == 0 {
c.RAG.ChunkSize = 500
}
if c.RAG.ChunkOverlap == 0 {
c.RAG.ChunkOverlap = 50
}
if c.RAG.TopK == 0 {
c.RAG.TopK = 5
}
if c.RAG.Tokenize == "" {
c.RAG.Tokenize = "unicode61"
}
return nil
}
func (c *Config) DefaultProvider() *Provider {
for i := range c.Providers {
if c.Providers[i].Default {
return &c.Providers[i]
}
}
return &c.Providers[0]
}
func (c *Config) Addr() string {
return fmt.Sprintf("%s:%d", c.Server.Host, c.Server.Port)
}
func splitPath(path string) []string {
return strings.Split(path, string(os.PathSeparator))
}