From 82b11caa42acd05ea111b77850b7798937c6d49f Mon Sep 17 00:00:00 2001 From: Max Date: Fri, 22 Aug 2025 11:18:04 +0800 Subject: [PATCH] Implement hit and vote management enhancements in API - Added error handling for uninitialized knowledge base in hit and vote endpoints. - Implemented ScrollHits and ScrollVotes methods to retrieve paginated results with filtering options. - Updated AddHits and AddVotes functions to handle request validation and response formatting. - Enhanced response structures to include updated counts for added and removed hits and votes. - Introduced new request types for updating hits and votes, allowing for default reactions in segments. --- openapi/kb/hit.go | 201 +++++++++++++++++++++++++++++++----------- openapi/kb/kb.go | 2 - openapi/kb/score.go | 16 +++- openapi/kb/types.go | 9 +- openapi/kb/vote.go | 204 ++++++++++++++++++++++++++++++++----------- openapi/kb/weight.go | 4 +- 6 files changed, 329 insertions(+), 107 deletions(-) diff --git a/openapi/kb/hit.go b/openapi/kb/hit.go index 27c68491..71ab7139 100644 --- a/openapi/kb/hit.go +++ b/openapi/kb/hit.go @@ -6,6 +6,8 @@ import ( "strings" "github.com/gin-gonic/gin" + "github.com/yaoapp/gou/graphrag/types" + "github.com/yaoapp/yao/kb" "github.com/yaoapp/yao/openapi/response" ) @@ -79,17 +81,47 @@ func ScrollHits(c *gin.Context) { options["filter"] = filter } - // TODO: Implement document permission validation for docID - // TODO: Implement scroll hits logic with GraphRag or database - // TODO: Call kb.Instance.ScrollHits(c.Request.Context(), segmentID, options) + // Check if kb.Instance is available + if kb.Instance == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Knowledge base not initialized", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } - // Return mock response for now - result := gin.H{ - "hits": []interface{}{}, - "scroll_id": nil, - "has_more": false, - "total": 0, - "options": options, + // Convert options to ScrollHitsOptions + scrollOptions := &types.ScrollHitsOptions{ + SegmentID: segmentID, + Limit: options["limit"].(int), + } + + // Set cursor if provided + if scrollID, exists := options["scroll_id"]; exists && scrollID != nil { + scrollOptions.Cursor = scrollID.(string) + } + + // Set filters if provided + if filter, exists := options["filter"]; exists && filter != nil { + filterMap := filter.(map[string]interface{}) + if source, ok := filterMap["source"]; ok { + scrollOptions.Source = source.(string) + } + if scenario, ok := filterMap["scenario"]; ok { + scrollOptions.Scenario = scenario.(string) + } + } + + // Call GraphRag ScrollHits method + result, err := kb.Instance.ScrollHits(c.Request.Context(), docID, scrollOptions) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to scroll hits: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return } response.RespondWithSuccess(c, response.StatusOK, result) @@ -131,34 +163,12 @@ func GetHits(c *gin.Context) { filter["session_id"] = sessionID } - // Parse limit parameter (optional, for basic limiting without pagination) - var limit int - if limitStr := c.Query("limit"); limitStr != "" { - if l, err := strconv.Atoi(limitStr); err == nil && l > 0 { - limit = l - } + // TODO: Search functionality not implemented yet - reserved for future use + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Search hits functionality is reserved but not implemented yet", } - - // TODO: Implement document permission validation for docID - // TODO: Implement get hits logic (simple query without pagination) - // TODO: Call kb.Instance.GetHits(c.Request.Context(), segmentID, filter, limit) - - // Return mock response for now - result := gin.H{ - "hits": []interface{}{}, - "doc_id": docID, - "segment_id": segmentID, - "total": 0, - } - - if len(filter) > 0 { - result["filter"] = filter - } - if limit > 0 { - result["limit"] = limit - } - - response.RespondWithSuccess(c, response.StatusOK, result) + response.RespondWithError(c, response.StatusNotImplemented, errorResp) } // GetHit gets a specific hit by ID @@ -206,7 +216,7 @@ func GetHit(c *gin.Context) { }) } -// AddHits adds new hits to a segment +// AddHits adds new hits to a segment using UpdateHits implementation func AddHits(c *gin.Context) { // Extract docID from URL path docID := c.Param("docID") @@ -230,14 +240,80 @@ func AddHits(c *gin.Context) { return } - // TODO: Implement document permission validation for docID - // TODO: Implement add hit logic - c.JSON(http.StatusOK, gin.H{ - "message": "Hit added successfully", - "doc_id": docID, - "segment_id": segmentID, - "hit_id": "placeholder-hit-id", - }) + // Parse request body for hit data + var req UpdateHitRequest + if err := c.ShouldBindJSON(&req); err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invalid request format: " + err.Error(), + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Validate request + if len(req.Segments) == 0 { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "At least one hit is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Ensure all hits are for the correct segment + for i := range req.Segments { + req.Segments[i].ID = segmentID + } + + // Check if kb.Instance is available + if kb.Instance == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Knowledge base not initialized", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Build options with default reaction from payload or create basic fallback + var options types.UpdateHitOptions + if req.DefaultReaction != nil { + // Use the default reaction provided in the request + options.Reaction = req.DefaultReaction + } else { + // Create basic fallback context for segments that don't have reaction + options.Reaction = &types.SegmentReaction{ + Source: "api", + Scenario: "hit", + Context: map[string]interface{}{ + "method": c.Request.Method, + "path": c.Request.URL.Path, + "client_ip": c.ClientIP(), + }, + } + } + + // Call GraphRag UpdateHits method + updatedCount, err := kb.Instance.UpdateHits(c.Request.Context(), docID, req.Segments, options) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to add hits: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + result := gin.H{ + "message": "Hits added successfully", + "doc_id": docID, + "segment_id": segmentID, + "hits": req.Segments, + "updated_count": updatedCount, + } + + response.RespondWithSuccess(c, response.StatusOK, result) } // UpdateHits updates hits in batch @@ -347,15 +423,42 @@ func RemoveHits(c *gin.Context) { return } - // TODO: Implement document permission validation for docID - // TODO: Implement batch remove hit logic + // Check if kb.Instance is available + if kb.Instance == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Knowledge base not initialized", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Build HitRemoval structs + var hitRemovals []types.HitRemoval + for _, hitID := range validHitIDs { + hitRemovals = append(hitRemovals, types.HitRemoval{ + SegmentID: segmentID, + HitID: hitID, + }) + } + + // Call GraphRag RemoveHits method + removedCount, err := kb.Instance.RemoveHits(c.Request.Context(), docID, hitRemovals) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to remove hits: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } result := gin.H{ "message": "Hits removed successfully", "doc_id": docID, "segment_id": segmentID, "hit_ids": validHitIDs, - "removed_count": len(validHitIDs), + "removed_count": removedCount, } response.RespondWithSuccess(c, response.StatusOK, result) diff --git a/openapi/kb/kb.go b/openapi/kb/kb.go index b03455f4..25d04b9f 100644 --- a/openapi/kb/kb.go +++ b/openapi/kb/kb.go @@ -62,7 +62,6 @@ func Attach(group *gin.RouterGroup, oauth types.OAuth) { group.GET("/documents/:docID/segments/:segmentID/votes/search", GetVotes) group.GET("/documents/:docID/segments/:segmentID/votes/:voteID", GetVote) group.POST("/documents/:docID/segments/:segmentID/votes", AddVotes) - group.PUT("/documents/:docID/segments/:segmentID/votes", UpdateVotes) group.DELETE("/documents/:docID/segments/:segmentID/votes", RemoveVotes) // Segment hits management @@ -70,7 +69,6 @@ func Attach(group *gin.RouterGroup, oauth types.OAuth) { group.GET("/documents/:docID/segments/:segmentID/hits/search", GetHits) group.GET("/documents/:docID/segments/:segmentID/hits/:hitID", GetHit) group.POST("/documents/:docID/segments/:segmentID/hits", AddHits) - group.PUT("/documents/:docID/segments/:segmentID/hits", UpdateHits) group.DELETE("/documents/:docID/segments/:segmentID/hits", RemoveHits) // Search Management diff --git a/openapi/kb/score.go b/openapi/kb/score.go index a76eb422..f54d9d82 100644 --- a/openapi/kb/score.go +++ b/openapi/kb/score.go @@ -5,6 +5,7 @@ import ( "strings" "github.com/gin-gonic/gin" + "github.com/yaoapp/yao/kb" "github.com/yaoapp/yao/openapi/response" ) @@ -64,15 +65,22 @@ func UpdateScores(c *gin.Context) { } } - // TODO: Implement document permission validation for docID - // TODO: Implement batch update scores logic - // TODO: Call kb.Instance.UpdateScores(c.Request.Context(), docID, req.Scores) + // Call GraphRag UpdateScores method (without Compute option) + updatedCount, err := kb.Instance.UpdateScores(c.Request.Context(), docID, req.Scores) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to update scores: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } result := gin.H{ "message": "Scores updated successfully", "doc_id": docID, "scores": req.Scores, - "updated_count": len(req.Scores), + "updated_count": updatedCount, } response.RespondWithSuccess(c, response.StatusOK, result) diff --git a/openapi/kb/types.go b/openapi/kb/types.go index 783d5f4b..22224adb 100644 --- a/openapi/kb/types.go +++ b/openapi/kb/types.go @@ -131,7 +131,14 @@ type UpdateSegmentsRequest struct { // UpdateVoteRequest represents the request for UpdateVote API type UpdateVoteRequest struct { - Segments []types.SegmentVote `json:"segments" binding:"required"` + Segments []types.SegmentVote `json:"segments" binding:"required"` + DefaultReaction *types.SegmentReaction `json:"default_reaction,omitempty"` // Optional default context for segments that don't have reaction +} + +// UpdateHitRequest represents the request for UpdateHit API +type UpdateHitRequest struct { + Segments []types.SegmentHit `json:"segments" binding:"required"` + DefaultReaction *types.SegmentReaction `json:"default_reaction,omitempty"` // Optional default context for segments that don't have reaction } // UpdateScoreRequest represents the request for UpdateScore API diff --git a/openapi/kb/vote.go b/openapi/kb/vote.go index 23486122..2859e421 100644 --- a/openapi/kb/vote.go +++ b/openapi/kb/vote.go @@ -6,6 +6,8 @@ import ( "strings" "github.com/gin-gonic/gin" + "github.com/yaoapp/gou/graphrag/types" + "github.com/yaoapp/yao/kb" "github.com/yaoapp/yao/openapi/response" ) @@ -76,17 +78,50 @@ func ScrollVotes(c *gin.Context) { options["filter"] = filter } - // TODO: Implement document permission validation for docID - // TODO: Implement scroll votes logic with GraphRag or database - // TODO: Call kb.Instance.ScrollVotes(c.Request.Context(), segmentID, options) + // Check if kb.Instance is available + if kb.Instance == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Knowledge base not initialized", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } - // Return mock response for now - result := gin.H{ - "votes": []interface{}{}, - "scroll_id": nil, - "has_more": false, - "total": 0, - "options": options, + // Convert options to ScrollVotesOptions + scrollOptions := &types.ScrollVotesOptions{ + SegmentID: segmentID, + Limit: options["limit"].(int), + } + + // Set cursor if provided + if scrollID, exists := options["scroll_id"]; exists && scrollID != nil { + scrollOptions.Cursor = scrollID.(string) + } + + // Set filters if provided + if filter, exists := options["filter"]; exists && filter != nil { + filterMap := filter.(map[string]interface{}) + if voteType, ok := filterMap["vote_type"]; ok { + scrollOptions.VoteType = types.VoteType(voteType.(string)) + } + if source, ok := filterMap["source"]; ok { + scrollOptions.Source = source.(string) + } + if scenario, ok := filterMap["scenario"]; ok { + scrollOptions.Scenario = scenario.(string) + } + } + + // Call GraphRag ScrollVotes method + result, err := kb.Instance.ScrollVotes(c.Request.Context(), docID, scrollOptions) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to scroll votes: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return } response.RespondWithSuccess(c, response.StatusOK, result) @@ -125,34 +160,12 @@ func GetVotes(c *gin.Context) { filter["user_id"] = userID } - // Parse limit parameter (optional, for basic limiting without pagination) - var limit int - if limitStr := c.Query("limit"); limitStr != "" { - if l, err := strconv.Atoi(limitStr); err == nil && l > 0 { - limit = l - } + // TODO: Search functionality not implemented yet - reserved for future use + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Search votes functionality is reserved but not implemented yet", } - - // TODO: Implement document permission validation for docID - // TODO: Implement get votes logic (simple query without pagination) - // TODO: Call kb.Instance.GetVotes(c.Request.Context(), segmentID, filter, limit) - - // Return mock response for now - result := gin.H{ - "votes": []interface{}{}, - "doc_id": docID, - "segment_id": segmentID, - "total": 0, - } - - if len(filter) > 0 { - result["filter"] = filter - } - if limit > 0 { - result["limit"] = limit - } - - response.RespondWithSuccess(c, response.StatusOK, result) + response.RespondWithError(c, response.StatusNotImplemented, errorResp) } // GetVote gets a specific vote by ID @@ -200,7 +213,7 @@ func GetVote(c *gin.Context) { }) } -// AddVotes adds new votes to a segment +// AddVotes adds new votes to a segment using UpdateVotes implementation func AddVotes(c *gin.Context) { // Extract docID from URL path docID := c.Param("docID") @@ -224,14 +237,80 @@ func AddVotes(c *gin.Context) { return } - // TODO: Implement document permission validation for docID - // TODO: Implement add vote logic - c.JSON(http.StatusOK, gin.H{ - "message": "Vote added successfully", - "doc_id": docID, - "segment_id": segmentID, - "vote_id": "placeholder-vote-id", - }) + // Parse request body for vote data + var req UpdateVoteRequest + if err := c.ShouldBindJSON(&req); err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "Invalid request format: " + err.Error(), + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Validate request + if len(req.Segments) == 0 { + errorResp := &response.ErrorResponse{ + Code: response.ErrInvalidRequest.Code, + ErrorDescription: "At least one vote is required", + } + response.RespondWithError(c, response.StatusBadRequest, errorResp) + return + } + + // Ensure all votes are for the correct segment + for i := range req.Segments { + req.Segments[i].ID = segmentID + } + + // Check if kb.Instance is available + if kb.Instance == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Knowledge base not initialized", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Build options with default reaction from payload or create basic fallback + var options types.UpdateVoteOptions + if req.DefaultReaction != nil { + // Use the default reaction provided in the request + options.Reaction = req.DefaultReaction + } else { + // Create basic fallback context for segments that don't have reaction + options.Reaction = &types.SegmentReaction{ + Source: "api", + Scenario: "vote", + Context: map[string]interface{}{ + "method": c.Request.Method, + "path": c.Request.URL.Path, + "client_ip": c.ClientIP(), + }, + } + } + + // Call GraphRag UpdateVotes method + updatedCount, err := kb.Instance.UpdateVotes(c.Request.Context(), docID, req.Segments, options) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to add votes: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + result := gin.H{ + "message": "Votes added successfully", + "doc_id": docID, + "segment_id": segmentID, + "votes": req.Segments, + "updated_count": updatedCount, + } + + response.RespondWithSuccess(c, response.StatusOK, result) } // UpdateVotes updates votes in batch @@ -341,15 +420,42 @@ func RemoveVotes(c *gin.Context) { return } - // TODO: Implement document permission validation for docID - // TODO: Implement batch remove vote logic + // Check if kb.Instance is available + if kb.Instance == nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Knowledge base not initialized", + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } + + // Build VoteRemoval structs + var voteRemovals []types.VoteRemoval + for _, voteID := range validVoteIDs { + voteRemovals = append(voteRemovals, types.VoteRemoval{ + SegmentID: segmentID, + VoteID: voteID, + }) + } + + // Call GraphRag RemoveVotes method + removedCount, err := kb.Instance.RemoveVotes(c.Request.Context(), docID, voteRemovals) + if err != nil { + errorResp := &response.ErrorResponse{ + Code: response.ErrServerError.Code, + ErrorDescription: "Failed to remove votes: " + err.Error(), + } + response.RespondWithError(c, response.StatusInternalServerError, errorResp) + return + } result := gin.H{ "message": "Votes removed successfully", "doc_id": docID, "segment_id": segmentID, "vote_ids": validVoteIDs, - "removed_count": len(validVoteIDs), + "removed_count": removedCount, } response.RespondWithSuccess(c, response.StatusOK, result) diff --git a/openapi/kb/weight.go b/openapi/kb/weight.go index 1b7d5ac3..866c003c 100644 --- a/openapi/kb/weight.go +++ b/openapi/kb/weight.go @@ -55,8 +55,8 @@ func UpdateWeights(c *gin.Context) { // TODO: Implement document permission validation for docID - // Perform batch update weight operation - updatedCount, err := kb.Instance.UpdateWeight(c.Request.Context(), docID, req.Weights) + // Call GraphRag UpdateWeights method (without Compute option) + updatedCount, err := kb.Instance.UpdateWeights(c.Request.Context(), docID, req.Weights) if err != nil { errorResp := &response.ErrorResponse{ Code: response.ErrServerError.Code,