Skip to content

Commit 189d1fd

Browse files
committed
🎨 Improve AI configuration normalization
1 parent 3d017d3 commit 189d1fd

3 files changed

Lines changed: 65 additions & 97 deletions

File tree

kernel/api/setting.go

Lines changed: 1 addition & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -198,57 +198,12 @@ func setAI(c *gin.Context) {
198198
return
199199
}
200200

201-
for _, p := range ai.Providers {
202-
if nil == p {
203-
continue
204-
}
205-
if 1 > p.RequestTimeout {
206-
p.RequestTimeout = 30
207-
}
208-
}
209-
210-
if nil != ai.Editing {
211-
if 0 > ai.Editing.MaxCompletionTokens {
212-
ai.Editing.MaxCompletionTokens = 0
213-
}
214-
if 0 > ai.Editing.Temperature || 2 < ai.Editing.Temperature {
215-
ai.Editing.Temperature = 1.0
216-
}
217-
if 1 > ai.Editing.MaxHistoryMessages || 64 < ai.Editing.MaxHistoryMessages {
218-
ai.Editing.MaxHistoryMessages = 7
219-
}
220-
}
221-
222-
if len(ai.Providers) == 0 {
223-
ai.Providers = model.Conf.AI.Providers
224-
}
225-
if nil == ai.MCP {
226-
ai.MCP = model.Conf.AI.MCP
227-
}
228-
if nil == ai.Embedding {
229-
ai.Embedding = model.Conf.AI.Embedding
230-
}
231-
if nil == ai.Agent {
232-
ai.Agent = model.Conf.AI.Agent
233-
}
234-
if nil == ai.Editing {
235-
ai.Editing = model.Conf.AI.Editing
236-
}
237-
238-
for i, p := range ai.Providers {
239-
if nil == p {
240-
continue
241-
}
242-
if "" == p.ID && i < len(model.Conf.AI.Providers) && nil != model.Conf.AI.Providers[i] {
243-
p.ID = model.Conf.AI.Providers[i].ID
244-
}
245-
}
246201
model.Conf.AI = ai
247202

248203
model.Conf.AI.Normalize()
249204
model.Conf.Save()
250205

251-
ret.Data = ai
206+
ret.Data = model.Conf.AI
252207
}
253208

254209
func setFlashcard(c *gin.Context) {

kernel/conf/ai.go

Lines changed: 63 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@ import (
2121
"encoding/json"
2222
"os"
2323
"strconv"
24+
"strings"
2425

2526
"github.com/88250/lute/ast"
2627
"github.com/siyuan-note/siyuan/kernel/util"
@@ -102,24 +103,32 @@ func defaultEmbedding() *Embedding {
102103
return &Embedding{Timeout: 30}
103104
}
104105

106+
func defaultAgent() *Agent {
107+
return &Agent{
108+
SessionTimeout: 600,
109+
ConfirmTimeout: 120,
110+
MaxRetries: 3,
111+
Temperature: 1.0,
112+
MaxCompletionTokens: 0,
113+
MaxToolCallRounds: 64,
114+
}
115+
}
116+
117+
func defaultEditing() *Editing {
118+
return &Editing{
119+
MaxHistoryMessages: 7,
120+
Temperature: 1.0,
121+
MaxCompletionTokens: 0,
122+
}
123+
}
124+
105125
func NewAI() *AI {
106126
ai := &AI{
107127
Providers: []*Provider{},
108128
MCP: &MCP{Servers: []MCPServer{}},
109129
Embedding: defaultEmbedding(),
110-
Agent: &Agent{
111-
SessionTimeout: 600,
112-
ConfirmTimeout: 120,
113-
MaxRetries: 3,
114-
Temperature: 1.0,
115-
MaxCompletionTokens: 0,
116-
MaxToolCallRounds: 64,
117-
},
118-
Editing: &Editing{
119-
MaxHistoryMessages: 7,
120-
Temperature: 1.0,
121-
MaxCompletionTokens: 0,
122-
},
130+
Agent: defaultAgent(),
131+
Editing: defaultEditing(),
123132
}
124133

125134
apiKey := os.Getenv("SIYUAN_OPENAI_API_KEY")
@@ -286,25 +295,64 @@ func (ai *AI) Normalize() {
286295
} else if ai.MCP.Servers == nil {
287296
ai.MCP.Servers = []MCPServer{}
288297
}
298+
if ai.Agent == nil {
299+
ai.Agent = defaultAgent()
300+
}
301+
if ai.Editing == nil {
302+
ai.Editing = defaultEditing()
303+
} else {
304+
if 0 > ai.Editing.MaxCompletionTokens {
305+
ai.Editing.MaxCompletionTokens = 0
306+
}
307+
if 0 > ai.Editing.Temperature {
308+
ai.Editing.Temperature = 0
309+
} else if 2 < ai.Editing.Temperature {
310+
ai.Editing.Temperature = 2
311+
}
312+
if 1 > ai.Editing.MaxHistoryMessages {
313+
ai.Editing.MaxHistoryMessages = 1
314+
} else if 64 < ai.Editing.MaxHistoryMessages {
315+
ai.Editing.MaxHistoryMessages = 64
316+
}
317+
}
318+
providers := make([]*Provider, 0, len(ai.Providers))
289319
for _, p := range ai.Providers {
290320
if p == nil {
291321
continue
292322
}
293-
if p.Models == nil {
294-
p.Models = []*Model{}
323+
p.BaseURL = strings.TrimSpace(p.BaseURL)
324+
if "" == p.BaseURL {
325+
p.BaseURL = "https://api.openai.com/v1"
326+
}
327+
p.DisplayName = strings.TrimSpace(p.DisplayName)
328+
p.APIKey = strings.TrimSpace(p.APIKey)
329+
if 1 > p.RequestTimeout {
330+
p.RequestTimeout = 30
331+
} else if 600 < p.RequestTimeout {
332+
p.RequestTimeout = 600
295333
}
296334
if !ast.IsNodeIDPattern(p.ID) {
297335
p.ID = ast.NewNodeID()
298336
}
337+
models := make([]*Model, 0, len(p.Models))
299338
for _, m := range p.Models {
300339
if m == nil {
301340
continue
302341
}
342+
m.Name = strings.TrimSpace(m.Name)
343+
if "" == m.Name {
344+
m.Name = "model"
345+
}
346+
m.DisplayName = strings.TrimSpace(m.DisplayName)
303347
if !ast.IsNodeIDPattern(m.ID) {
304348
m.ID = ast.NewNodeID()
305349
}
350+
models = append(models, m)
306351
}
352+
p.Models = models
353+
providers = append(providers, p)
307354
}
355+
ai.Providers = providers
308356
if ai.Embedding == nil {
309357
ai.Embedding = defaultEmbedding()
310358
}

kernel/model/conf.go

Lines changed: 1 addition & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -591,40 +591,7 @@ func InitConf() {
591591
if nil == Conf.AI {
592592
Conf.AI = conf.NewAI()
593593
}
594-
if nil == Conf.AI.Agent {
595-
Conf.AI.Agent = &conf.Agent{
596-
SessionTimeout: 600,
597-
ConfirmTimeout: 120,
598-
MaxRetries: 3,
599-
Temperature: 1.0,
600-
MaxCompletionTokens: 0,
601-
MaxToolCallRounds: 64,
602-
}
603-
}
604-
if nil == Conf.AI.Editing {
605-
Conf.AI.Editing = &conf.Editing{
606-
MaxHistoryMessages: 7,
607-
Temperature: 1.0,
608-
MaxCompletionTokens: 0,
609-
}
610-
}
611-
for _, p := range Conf.AI.Providers {
612-
if nil == p {
613-
continue
614-
}
615-
if 1 > p.RequestTimeout {
616-
p.RequestTimeout = 30
617-
}
618-
}
619-
if 0 > Conf.AI.Editing.MaxCompletionTokens {
620-
Conf.AI.Editing.MaxCompletionTokens = 0
621-
}
622-
if 0 > Conf.AI.Editing.Temperature || 2 < Conf.AI.Editing.Temperature {
623-
Conf.AI.Editing.Temperature = 1.0
624-
}
625-
if 1 > Conf.AI.Editing.MaxHistoryMessages || 64 < Conf.AI.Editing.MaxHistoryMessages {
626-
Conf.AI.Editing.MaxHistoryMessages = 7
627-
}
594+
Conf.AI.Normalize()
628595

629596
for _, p := range Conf.AI.Providers {
630597
if p == nil || len(p.APIKey) == 0 {
@@ -658,8 +625,6 @@ func InitConf() {
658625
Conf.AI.Embedding.Name)
659626
}
660627

661-
Conf.AI.Normalize()
662-
663628
Conf.ReadOnly = util.ReadOnly
664629

665630
if "" != util.AccessAuthCode {

0 commit comments

Comments
 (0)