Sanitize generated DSL and update execution flow in search handler
- Introduce a new sanitizeDSL function to clean up unsupported wildcards in the generated DSL. - Adjust the sequence of operations in SearchWithContext to reflect the new sanitization step. - Update comments to accurately reflect the changes in the execution flow and improve code clarity.
This commit is contained in:
parent
efb612734d
commit
3fd27627a1
1 changed files with 37 additions and 13 deletions
|
|
@ -164,10 +164,13 @@ func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Reques
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// 3. Merge preset conditions into generated DSL
|
// 3. Sanitize generated DSL (remove unsupported wildcards like "*")
|
||||||
|
h.sanitizeDSL(result.DSL)
|
||||||
|
|
||||||
|
// 4. Merge preset conditions into generated DSL
|
||||||
h.mergeDSLConditions(result.DSL, req)
|
h.mergeDSLConditions(result.DSL, req)
|
||||||
|
|
||||||
// 4. Execute QueryDSL using gou query engine
|
// 5. Execute QueryDSL using gou query engine
|
||||||
records, err := h.executeDSL(result.DSL)
|
records, err := h.executeDSL(result.DSL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return &types.Result{
|
return &types.Result{
|
||||||
|
|
@ -181,7 +184,7 @@ func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Reques
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// 5. Determine the primary model for result formatting
|
// 6. Determine the primary model for result formatting
|
||||||
// Use the "from" table from DSL, or first model
|
// Use the "from" table from DSL, or first model
|
||||||
primaryModelID := modelIDs[0]
|
primaryModelID := modelIDs[0]
|
||||||
if result.DSL.From != nil && result.DSL.From.Name != "" {
|
if result.DSL.From != nil && result.DSL.From.Name != "" {
|
||||||
|
|
@ -199,7 +202,7 @@ func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Reques
|
||||||
primaryModel, _ = model.Get(primaryModelID) // May be nil, that's ok
|
primaryModel, _ = model.Get(primaryModelID) // May be nil, that's ok
|
||||||
}
|
}
|
||||||
|
|
||||||
// 6. Convert records to ResultItems
|
// 7. Convert records to ResultItems
|
||||||
items := h.convertToResultItems(records, primaryModelID, primaryModel, req.Source)
|
items := h.convertToResultItems(records, primaryModelID, primaryModel, req.Source)
|
||||||
|
|
||||||
// Apply limit
|
// Apply limit
|
||||||
|
|
@ -207,7 +210,7 @@ func (h *Handler) SearchWithContext(ctx *agentContext.Context, req *types.Reques
|
||||||
items = items[:maxResults]
|
items = items[:maxResults]
|
||||||
}
|
}
|
||||||
|
|
||||||
// 7. Convert DSL to map for storage
|
// 8. Convert DSL to map for storage
|
||||||
dslMap := h.dslToMap(result.DSL)
|
dslMap := h.dslToMap(result.DSL)
|
||||||
|
|
||||||
return &types.Result{
|
return &types.Result{
|
||||||
|
|
@ -282,31 +285,52 @@ func (h *Handler) buildModelSchema(mod *model.Model) map[string]interface{} {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// executeDSL executes the QueryDSL and returns records
|
// sanitizeDSL cleans up LLM-generated DSL to remove unsupported constructs.
|
||||||
func (h *Handler) executeDSL(dsl interface{}) ([]map[string]interface{}, error) {
|
// The QueryDSL engine does not support wildcard "*" in select fields;
|
||||||
// Get the default query engine
|
// an empty select list naturally returns all columns.
|
||||||
|
func (h *Handler) sanitizeDSL(dsl *gou.QueryDSL) {
|
||||||
|
if dsl == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(dsl.Select) > 0 {
|
||||||
|
cleaned := make([]gou.Expression, 0, len(dsl.Select))
|
||||||
|
for _, expr := range dsl.Select {
|
||||||
|
if expr.Field != "*" {
|
||||||
|
cleaned = append(cleaned, expr)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
dsl.Select = cleaned
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// executeDSL executes the QueryDSL and returns records.
|
||||||
|
// Uses recover to convert panics from MustGet into errors.
|
||||||
|
func (h *Handler) executeDSL(dsl interface{}) (records []map[string]interface{}, err error) {
|
||||||
|
defer func() {
|
||||||
|
if r := recover(); r != nil {
|
||||||
|
err = fmt.Errorf("query execution panic: %v", r)
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
engine, err := query.Select("default")
|
engine, err := query.Select("default")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("query engine not found: %w", err)
|
return nil, fmt.Errorf("query engine not found: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Marshal DSL to JSON
|
|
||||||
dslJSON, err := json.Marshal(dsl)
|
dslJSON, err := json.Marshal(dsl)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to marshal DSL: %w", err)
|
return nil, fmt.Errorf("failed to marshal DSL: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Load and execute the query
|
|
||||||
q, err := engine.Load(json.RawMessage(dslJSON))
|
q, err := engine.Load(json.RawMessage(dslJSON))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to load DSL: %w", err)
|
return nil, fmt.Errorf("failed to load DSL: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Execute query
|
|
||||||
rawRecords := q.Get(nil)
|
rawRecords := q.Get(nil)
|
||||||
|
|
||||||
// Convert to map[string]interface{}
|
records = make([]map[string]interface{}, 0, len(rawRecords))
|
||||||
records := make([]map[string]interface{}, 0, len(rawRecords))
|
|
||||||
for _, rec := range rawRecords {
|
for _, rec := range rawRecords {
|
||||||
records = append(records, map[string]interface{}(rec))
|
records = append(records, map[string]interface{}(rec))
|
||||||
}
|
}
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue