Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions admin/server/ai.go
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,7 @@ func (s *Server) Complete(ctx context.Context, req *adminv1.CompleteRequest) (*a
Messages: messages,
Tools: req.Tools,
OutputSchema: outputSchema,
CacheKey: req.CacheKey,
})
if err != nil {
return nil, err
Expand Down
5 changes: 5 additions & 0 deletions proto/gen/rill/admin/v1/admin.swagger.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -5171,6 +5171,11 @@ definitions:
outputJsonSchema:
type: string
title: Optional output JSON schema
cacheKey:
type: string
description: |-
Optional key identifying a series of requests that share a prompt prefix (e.g. an AI session ID).
Providers may use it to improve prompt cache routing.
v1CompleteResponse:
type: object
properties:
Expand Down
82 changes: 47 additions & 35 deletions proto/gen/rill/admin/v1/ai.pb.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

2 changes: 2 additions & 0 deletions proto/gen/rill/admin/v1/ai.pb.validate.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

5 changes: 5 additions & 0 deletions proto/gen/rill/admin/v1/openapi.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -382,6 +382,11 @@ components:
type: object
v1CompleteRequest:
properties:
cacheKey:
description: |-
Optional key identifying a series of requests that share a prompt prefix (e.g. an AI session ID).
Providers may use it to improve prompt cache routing.
type: string
messages:
description: Input message(s) for the AI to complete.
items:
Expand Down
5 changes: 5 additions & 0 deletions proto/gen/rill/admin/v1/public.openapi.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -382,6 +382,11 @@ components:
type: object
v1CompleteRequest:
properties:
cacheKey:
description: |-
Optional key identifying a series of requests that share a prompt prefix (e.g. an AI session ID).
Providers may use it to improve prompt cache routing.
type: string
messages:
description: Input message(s) for the AI to complete.
items:
Expand Down
3 changes: 3 additions & 0 deletions proto/rill/admin/v1/ai.proto
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,9 @@ message CompleteRequest {
repeated rill.ai.v1.Tool tools = 2;
// Optional output JSON schema
string output_json_schema = 3;
// Optional key identifying a series of requests that share a prompt prefix (e.g. an AI session ID).
// Providers may use it to improve prompt cache routing.
string cache_key = 4;
}

message CompleteResponse {
Expand Down
80 changes: 67 additions & 13 deletions runtime/ai/ai.go
Original file line number Diff line number Diff line change
Expand Up @@ -29,13 +29,19 @@ import (
"go.uber.org/zap"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/proto"
"google.golang.org/protobuf/types/known/structpb"
)

// maxMessageSizeBytes is the maximum allowed size of a message's contents.
// Exceeding it will result in an error.
const maxMessageSizeBytes = 100 * 1024 // 100 KB

// reservedInputTokens is headroom subtracted from the model's input token limit to make room for
// tool definitions, model output (some providers share one context window between input and output),
// and token estimation error.
const reservedInputTokens = 64_000

// Tracer for instrumenting requests.
var tracer = otel.Tracer("github.com/rilldata/rill/runtime/ai")

Expand Down Expand Up @@ -1212,6 +1218,14 @@ func (s *Session) Complete(ctx context.Context, name string, out any, opts *Comp
}
defer release()

// Determine the token budget for input messages based on the model's input token limit.
// The reserved headroom is capped at a quarter of the limit so small limits keep a usable budget.
maxInputTokens := llm.MaxInputTokens()
if maxInputTokens <= 0 {
maxInputTokens = drivers.DefaultAIMaxInputTokens
}
messagesTokenBudget := maxInputTokens - min(reservedInputTokens, maxInputTokens/4)

// Setup input messages.
messages := slices.Clone(opts.Messages)

Expand Down Expand Up @@ -1287,7 +1301,7 @@ func (s *Session) Complete(ctx context.Context, name string, out any, opts *Comp
}

// Truncate messages to fit within LLM context window.
truncMessages := maybeTruncateMessages(messages)
truncMessages := maybeTruncateMessages(messages, messagesTokenBudget)

// Telemetry
iterations++
Expand All @@ -1304,6 +1318,7 @@ func (s *Session) Complete(ctx context.Context, name string, out any, opts *Comp
Messages: truncMessages,
Tools: tools,
OutputSchema: outputSchema,
CacheKey: s.id,
})
llmCancel()

Expand Down Expand Up @@ -1600,26 +1615,59 @@ func NewTextCompletionMessage(role Role, content string) *aiv1.CompletionMessage
}
}

// Tuning constants for maybeTruncateMessages.
const (
truncateKeepFirst = 4 // Always keep the first messages for context
truncateStep = 20 // Skip messages in multiples of this
)

// maybeTruncateMessages keeps recent messages and a few early ones for context.
// It's a simple placeholder strategy. In the future, we'll enhance this with AI summarization.
func maybeTruncateMessages(messages []*aiv1.CompletionMessage) []*aiv1.CompletionMessage {
const (
maxMessages = 20 // Keep up to 20 messages total
keepFirst = 4 // Always keep first 4 messages for context
keepLast = 16 // Keep last 16 messages
)
//
// Truncation triggers when the estimated token count of the messages exceeds maxTokens.
//
// LLM prompt caching requires the message prefix to be byte-stable across requests, so the number
// of skipped messages only changes in coarse steps of truncateStep: while up to truncateStep new
// messages accumulate, the result is an append-only extension of the previous result, causing a
// cache miss once per step instead of on every request.
func maybeTruncateMessages(messages []*aiv1.CompletionMessage, maxTokens int) []*aiv1.CompletionMessage {
if maxTokens <= 0 || len(messages) <= truncateKeepFirst+1 {
return messages
}

// Determine the smallest number of messages to skip such that the estimated token count of the
// kept messages fits within maxTokens. Skipping never reaches the last message: if the first
// and last messages alone exceed the budget, truncation can't help and we send them anyway.
var skipped int
var keptTokens int
est := make([]int, len(messages))
for i, m := range messages {
est[i] = estimateMessageTokens(m)
keptTokens += est[i]
}
for keptTokens > maxTokens && truncateKeepFirst+skipped < len(messages)-1 {
keptTokens -= est[truncateKeepFirst+skipped]
skipped++
}

if len(messages) <= maxMessages {
if skipped == 0 {
return messages
}

// Round the skipped count up to a multiple of truncateStep to keep the skipped range stable
// until a step's worth of additional tokens has accumulated.
skipped = ((skipped + truncateStep - 1) / truncateStep) * truncateStep
// do not skip
if skipped > len(messages)-truncateKeepFirst-1 {
skipped = len(messages) - truncateKeepFirst - 1
}

var result []*aiv1.CompletionMessage

// Keep first messages
result = append(result, messages[:keepFirst]...)
result = append(result, messages[:truncateKeepFirst]...)

// Add truncation indicator
skipped := len(messages) - keepFirst - keepLast
result = append(result, &aiv1.CompletionMessage{
Role: "system",
Content: []*aiv1.ContentBlock{
Expand All @@ -1631,9 +1679,8 @@ func maybeTruncateMessages(messages []*aiv1.CompletionMessage) []*aiv1.Completio
},
})

// Keep last messages
start := len(messages) - keepLast
result = append(result, messages[start:]...)
// Keep messages after the skipped range
result = append(result, messages[truncateKeepFirst+skipped:]...)

// Make sure there are no partial tool calls/results
unbalancedIDs := make(map[string]bool)
Expand All @@ -1660,6 +1707,13 @@ func maybeTruncateMessages(messages []*aiv1.CompletionMessage) []*aiv1.Completio
return result
}

// estimateMessageTokens conservatively estimates the number of LLM tokens in a message.
// It uses the proto-encoded size at 3 bytes per token: typical English text is ~4 bytes per token,
// but dense JSON and numeric content (common in tool results) can be as low as ~2-3.
func estimateMessageTokens(m *aiv1.CompletionMessage) int {
return proto.Size(m) / 3
}

// completionMessageID turns a UUID into a truncated ID suitable for use in completion messages (which don't require IDs to be globally unique).
func completionMessageID(id string) string {
return strings.ReplaceAll(id, "-", "")[0:16]
Expand Down
Loading
Loading