Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
Original file line number Diff line number Diff line change
Expand Up @@ -284,8 +284,10 @@ func buildModelCatalogFromModel(provider *Provider, model chat.Model) *ModelCata
"temperature",
"max_tokens",
"top_p",
"top_k",
"frequency_penalty",
"presence_penalty",
"repetition_penalty",
"stop",
"stream",
"n",
Expand Down
13 changes: 13 additions & 0 deletions services/llm-api/internal/domain/model/provider_model_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -209,6 +209,19 @@ func (s *ProviderModelService) BatchUpdateActive(ctx context.Context, filter Pro
return rowsAffected, nil
}

// FindCatalogByID returns a model catalog by ID (used for applying defaults).
func (s *ProviderModelService) FindCatalogByID(ctx context.Context, id uint) (*ModelCatalog, error) {
if id == 0 {
return nil, platformerrors.NewError(ctx, platformerrors.LayerDomain, platformerrors.ErrorTypeValidation, "model catalog ID is required", nil, "c9bde6a4-6bd1-4f8e-97df-e7b03c3f6f73")
}

catalog, err := s.modelCatalogRepo.FindByID(ctx, id)
if err != nil {
return nil, platformerrors.AsError(ctx, platformerrors.LayerDomain, err, "failed to find model catalog by ID")
}
return catalog, nil
}

func (s *ProviderModelService) BatchUpdateModelDisplayName(ctx context.Context, filter ProviderModelFilter, modelDisplayName string) (int64, error) {
rowsAffected, err := s.providerModelRepo.BatchUpdateModelDisplayName(ctx, filter, modelDisplayName)
if err != nil {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import (
"go.opentelemetry.io/otel/codes"

"jan-server/services/llm-api/internal/domain/conversation"
domainmodel "jan-server/services/llm-api/internal/domain/model"
"jan-server/services/llm-api/internal/domain/project"
"jan-server/services/llm-api/internal/domain/prompt"
"jan-server/services/llm-api/internal/domain/usersettings"
Expand All @@ -28,6 +29,8 @@ import (
"jan-server/services/llm-api/internal/utils/httpclients/chat"
"jan-server/services/llm-api/internal/utils/idgen"
"jan-server/services/llm-api/internal/utils/platformerrors"

"github.com/shopspring/decimal"
)

const ConversationReferrerContextKey = "conversation_referrer"
Expand Down Expand Up @@ -201,6 +204,16 @@ func (h *ChatHandler) CreateChatCompletion(
// Override the request model with the provider's original model ID
request.Model = selectedProviderModel.ProviderOriginalModelID

// Optionally load model catalog (used later to apply default parameters)
var modelCatalog *domainmodel.ModelCatalog
if selectedProviderModel.ModelCatalogID != nil {
modelCatalog, err = h.providerHandler.GetModelCatalogByID(ctx, *selectedProviderModel.ModelCatalogID)
if err != nil {
log := logger.GetLogger()
log.Warn().Err(err).Uint("model_catalog_id", *selectedProviderModel.ModelCatalogID).Msg("failed to load model catalog defaults")
}
}

// Resolve jan_* media placeholders (best-effort)
request.Messages = h.resolveMediaPlaceholders(ctx, reqCtx, request.Messages)

Expand Down Expand Up @@ -283,12 +296,20 @@ func (h *ChatHandler) CreateChatCompletion(
}

// Handle streaming vs non-streaming
llmRequest := chat.CompletionRequest{
ChatCompletionRequest: request.ChatCompletionRequest,
TopK: request.TopK,
RepetitionPenalty: request.RepetitionPenalty,
}
if modelCatalog != nil {
h.applyModelDefaultsFromCatalog(&llmRequest, modelCatalog)
}
observability.AddSpanEvent(ctx, "calling_llm")
llmStartTime := time.Now()
if request.Stream {
response, err = h.streamCompletion(ctx, reqCtx, chatClient, conv, request.ChatCompletionRequest)
response, err = h.streamCompletion(ctx, reqCtx, chatClient, conv, llmRequest)
} else {
response, err = h.callCompletion(ctx, chatClient, request.ChatCompletionRequest)
response, err = h.callCompletion(ctx, chatClient, llmRequest)
}
llmDuration := time.Since(llmStartTime)

Expand Down Expand Up @@ -419,7 +440,7 @@ func (h *ChatHandler) CreateChatCompletion(
func (h *ChatHandler) callCompletion(
ctx context.Context,
chatClient *chat.ChatCompletionClient,
request openai.ChatCompletionRequest,
request chat.CompletionRequest,
) (*openai.ChatCompletionResponse, error) {
chatCompletion, err := chatClient.CreateChatCompletion(ctx, "", request)
if err != nil {
Expand All @@ -435,7 +456,7 @@ func (h *ChatHandler) streamCompletion(
reqCtx *gin.Context,
chatClient *chat.ChatCompletionClient,
conv *conversation.Conversation,
request openai.ChatCompletionRequest,
request chat.CompletionRequest,
) (*openai.ChatCompletionResponse, error) {
// Create callback to send conversation data before [DONE]
var beforeDoneCallback chat.BeforeDoneCallback
Expand Down Expand Up @@ -525,6 +546,69 @@ func (h *ChatHandler) resolveMediaPlaceholders(ctx context.Context, reqCtx *gin.
return messages
}

// applyModelDefaultsFromCatalog fills in missing request parameters using defaults from the model catalog.
func (h *ChatHandler) applyModelDefaultsFromCatalog(req *chat.CompletionRequest, catalog *domainmodel.ModelCatalog) {
if req == nil || catalog == nil {
return
}

defaults := catalog.SupportedParameters.Default
if len(defaults) == 0 {
return
}

if req.Temperature == 0 {
if val, ok := decimalToFloat32(defaults["temperature"]); ok {
req.Temperature = val
}
}
if req.TopP == 0 {
if val, ok := decimalToFloat32(defaults["top_p"]); ok {
req.TopP = val
}
}
if req.PresencePenalty == 0 {
if val, ok := decimalToFloat32(defaults["presence_penalty"]); ok {
req.PresencePenalty = val
}
}
if req.FrequencyPenalty == 0 {
if val, ok := decimalToFloat32(defaults["frequency_penalty"]); ok {
req.FrequencyPenalty = val
}
}
if req.MaxTokens == 0 {
if val, ok := decimalToInt(defaults["max_tokens"]); ok {
req.MaxTokens = val
}
}
if req.TopK == nil || (req.TopK != nil && *req.TopK == 0) {

Copilot AI Dec 9, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The condition req.TopK == nil || (req.TopK != nil && *req.TopK == 0) is redundant. If req.TopK == nil is true, the second part of the OR will never be evaluated. Simplify to req.TopK == nil || *req.TopK == 0.

Suggested change
if req.TopK == nil || (req.TopK != nil && *req.TopK == 0) {
if req.TopK == nil || *req.TopK == 0 {

Copilot uses AI. Check for mistakes.
if val, ok := decimalToInt(defaults["top_k"]); ok {
req.TopK = &val
}
}
if req.RepetitionPenalty == nil || (req.RepetitionPenalty != nil && *req.RepetitionPenalty == 0) {

Copilot AI Dec 9, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The condition req.RepetitionPenalty == nil || (req.RepetitionPenalty != nil && *req.RepetitionPenalty == 0) is redundant. If req.RepetitionPenalty == nil is true, the second part of the OR will never be evaluated. Simplify to req.RepetitionPenalty == nil || *req.RepetitionPenalty == 0.

Suggested change
if req.RepetitionPenalty == nil || (req.RepetitionPenalty != nil && *req.RepetitionPenalty == 0) {
if req.RepetitionPenalty == nil || *req.RepetitionPenalty == 0 {

Copilot uses AI. Check for mistakes.
if val, ok := decimalToFloat32(defaults["repetition_penalty"]); ok {
req.RepetitionPenalty = &val
}
}
}

func decimalToFloat32(val *decimal.Decimal) (float32, bool) {
if val == nil {
return 0, false
}
f, _ := val.Float64()

Copilot AI Dec 9, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The error returned from Float64() conversion is being silently discarded with _. If the decimal-to-float conversion fails, this function will return 0 without indication. Consider checking and handling the conversion error, or at least document why it's safe to ignore.

Copilot uses AI. Check for mistakes.
return float32(f), true
}

func decimalToInt(val *decimal.Decimal) (int, bool) {
if val == nil {
return 0, false
}
return int(val.IntPart()), true
}

// getProjectInstruction loads the project instruction for the conversation, falling back to the stored snapshot.
func (h *ChatHandler) getProjectInstruction(ctx context.Context, userID uint, conv *conversation.Conversation) string {
if conv == nil || h.projectService == nil {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -257,6 +257,11 @@ func (h *ProviderHandler) UpdateProvider(
return &response, nil
}

// GetModelCatalogByID returns catalog details (used to apply default parameters).
func (providerHandler *ProviderHandler) GetModelCatalogByID(ctx context.Context, id uint) (*domainmodel.ModelCatalog, error) {
return providerHandler.providerModelService.FindCatalogByID(ctx, id)
}

func (h *ProviderHandler) DeleteProvider(ctx context.Context, publicID string) error {
if strings.TrimSpace(publicID) == "" {
return platformerrors.NewError(ctx, platformerrors.LayerHandler, platformerrors.ErrorTypeValidation, "provider public ID is required", nil, "0c3f68da-0aa4-4a7c-9cec-c22d47c86f8b")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,9 @@ import (
type ChatCompletionRequest struct {
openai.ChatCompletionRequest

TopK *int `json:"top_k,omitempty"`
RepetitionPenalty *float32 `json:"repetition_penalty,omitempty"`

// Conversation can be either a string (conversation ID) or a conversation object
// Items from this conversation are prepended to Messages for this response request.
// Input items and output items from this response are automatically added to this conversation after completion.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,13 @@ type ChatCompletionClient struct {
name string
}

// CompletionRequest extends the OpenAI chat request with provider-specific fields.
type CompletionRequest struct {
openai.ChatCompletionRequest
TopK *int `json:"top_k,omitempty"`
RepetitionPenalty *float32 `json:"repetition_penalty,omitempty"`
}

type functionCallAccumulator struct {
Name string
Arguments string
Expand All @@ -103,7 +110,7 @@ func NewChatCompletionClient(client *resty.Client, name, baseURL string) *ChatCo
}
}

func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey string, request openai.ChatCompletionRequest) (*openai.ChatCompletionResponse, error) {
func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey string, request CompletionRequest) (*openai.ChatCompletionResponse, error) {
// Start OpenTelemetry span for tracking
ctx, span := otel.Tracer("chat-completion-client").Start(ctx, "CreateChatCompletion",
trace.WithSpanKind(trace.SpanKindClient),
Expand All @@ -126,9 +133,41 @@ func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey
if request.TopP != 0 {
span.SetAttributes(attribute.Float64("llm.top_p", float64(request.TopP)))
}
if request.PresencePenalty != 0 {
span.SetAttributes(attribute.Float64("llm.presence_penalty", float64(request.PresencePenalty)))
}
if request.FrequencyPenalty != 0 {
span.SetAttributes(attribute.Float64("llm.frequency_penalty", float64(request.FrequencyPenalty)))
}

start := time.Now()

// Debug logging: Log request details before sending
log := logger.GetLogger()
topK := 0
if request.TopK != nil {
topK = *request.TopK
}
repetitionPenalty := float32(0)
if request.RepetitionPenalty != nil {
repetitionPenalty = *request.RepetitionPenalty
}
log.Info().
Str("provider", c.name).
Str("model", request.Model).
Int("messages", len(request.Messages)).
Bool("stream", request.Stream).
Float32("temperature", request.Temperature).
Int("max_tokens", request.MaxTokens).
Float32("top_p", request.TopP).
Int("top_k", topK).
Float32("presence_penalty", request.PresencePenalty).
Float32("frequency_penalty", request.FrequencyPenalty).
Float32("repetition_penalty", repetitionPenalty).
Msg("[ChatCompletion] Sending request to inference server")

log.Info().Interface("request", request).Msg("[ChatCompletion] Request body")

Copilot AI Dec 9, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Logging the entire request object at Info level (line 169) could expose sensitive data like API keys, user messages, or PII in logs. The request may contain conversation history with personal information. Consider using Debug level or sanitizing/redacting sensitive fields before logging.

Suggested change
log.Info().Interface("request", request).Msg("[ChatCompletion] Request body")
log.Debug().Interface("request", request).Msg("[ChatCompletion] Request body")

Copilot uses AI. Check for mistakes.

var respBody openai.ChatCompletionResponse
resp, err := c.prepareRequest(ctx, apiKey).
SetBody(request).
Expand Down Expand Up @@ -181,7 +220,7 @@ func (c *ChatCompletionClient) CreateChatCompletion(ctx context.Context, apiKey
return &respBody, nil
}

func (c *ChatCompletionClient) CreateChatCompletionStream(ctx context.Context, apiKey string, request openai.ChatCompletionRequest, opts ...StreamOption) (io.ReadCloser, error) {
func (c *ChatCompletionClient) CreateChatCompletionStream(ctx context.Context, apiKey string, request CompletionRequest, opts ...StreamOption) (io.ReadCloser, error) {
resp, err := c.doStreamingRequest(ctx, apiKey, request, opts...)
if err != nil {
return nil, err
Expand All @@ -208,10 +247,10 @@ func (c *ChatCompletionClient) CreateChatCompletionStream(ctx context.Context, a
}

func (c *ChatCompletionClient) StreamChatCompletionToContext(reqCtx *gin.Context, apiKey string, request openai.ChatCompletionRequest, opts ...StreamOption) (*openai.ChatCompletionResponse, error) {
return c.StreamChatCompletionToContextWithCallback(reqCtx, apiKey, request, nil, opts...)
return c.StreamChatCompletionToContextWithCallback(reqCtx, apiKey, CompletionRequest{ChatCompletionRequest: request}, nil, opts...)
}

func (c *ChatCompletionClient) StreamChatCompletionToContextWithCallback(reqCtx *gin.Context, apiKey string, request openai.ChatCompletionRequest, beforeDone BeforeDoneCallback, opts ...StreamOption) (*openai.ChatCompletionResponse, error) {
func (c *ChatCompletionClient) StreamChatCompletionToContextWithCallback(reqCtx *gin.Context, apiKey string, request CompletionRequest, beforeDone BeforeDoneCallback, opts ...StreamOption) (*openai.ChatCompletionResponse, error) {
// Start OpenTelemetry span for tracking streaming completion
ctx := reqCtx.Request.Context()
ctx, span := otel.Tracer("chat-completion-client").Start(ctx, "StreamChatCompletion",
Expand All @@ -235,9 +274,29 @@ func (c *ChatCompletionClient) StreamChatCompletionToContextWithCallback(reqCtx
if request.TopP != 0 {
span.SetAttributes(attribute.Float64("llm.top_p", float64(request.TopP)))
}
if request.PresencePenalty != 0 {
span.SetAttributes(attribute.Float64("llm.presence_penalty", float64(request.PresencePenalty)))
}
if request.FrequencyPenalty != 0 {
span.SetAttributes(attribute.Float64("llm.frequency_penalty", float64(request.FrequencyPenalty)))
}

start := time.Now()

// Debug logging: Log request details before sending
log := logger.GetLogger()
log.Debug().
Str("provider", c.name).
Str("model", request.Model).
Int("messages", len(request.Messages)).
Bool("stream", request.Stream).
Float32("temperature", request.Temperature).
Int("max_tokens", request.MaxTokens).
Float32("top_p", request.TopP).
Float32("presence_penalty", request.PresencePenalty).
Float32("frequency_penalty", request.FrequencyPenalty).
Msg("[StreamChatCompletion] Sending request to inference server")

// force to true to collect tokens
request.StreamOptions = &openai.StreamOptions{
IncludeUsage: true,
Expand Down Expand Up @@ -470,7 +529,33 @@ func (c *ChatCompletionClient) errorFromResponse(ctx context.Context, resp *rest
return platformerrors.NewError(ctx, platformerrors.LayerDomain, platformerrors.ErrorTypeExternal, fmt.Sprintf("%s: %s", message, trimmed), nil, "a1f46e0d-4017-4411-ac05-987946c3066d")
}

func (c *ChatCompletionClient) doStreamingRequest(ctx context.Context, apiKey string, request openai.ChatCompletionRequest, opts ...StreamOption) (*resty.Response, error) {
func (c *ChatCompletionClient) doStreamingRequest(ctx context.Context, apiKey string, request CompletionRequest, opts ...StreamOption) (*resty.Response, error) {
// Debug logging: Log request details for streaming calls
log := logger.GetLogger()
topK := 0
if request.TopK != nil {
topK = *request.TopK
}
repetitionPenalty := float32(0)
if request.RepetitionPenalty != nil {
repetitionPenalty = *request.RepetitionPenalty
}
log.Info().
Str("provider", c.name).
Str("model", request.Model).
Int("messages", len(request.Messages)).
Bool("stream", request.Stream).
Float32("temperature", request.Temperature).
Int("max_tokens", request.MaxTokens).
Float32("top_p", request.TopP).
Int("top_k", topK).
Float32("presence_penalty", request.PresencePenalty).
Float32("frequency_penalty", request.FrequencyPenalty).
Float32("repetition_penalty", repetitionPenalty).
Msg("[ChatCompletion][Stream] Sending request to inference server")

log.Info().Interface("request", request).Msg("[ChatCompletion][Stream] Request body")

Copilot AI Dec 9, 2025

Copy link

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Logging the entire request object at Info level (line 557) could expose sensitive data like API keys, user messages, or PII in logs. The request may contain conversation history with personal information. Consider using Debug level or sanitizing/redacting sensitive fields before logging.

Suggested change
log.Info().Interface("request", request).Msg("[ChatCompletion][Stream] Request body")
log.Debug().Interface("request", request).Msg("[ChatCompletion][Stream] Request body")

Copilot uses AI. Check for mistakes.

req := c.prepareRequest(ctx, apiKey).
SetBody(request).
SetDoNotParseResponse(true)
Expand Down Expand Up @@ -500,7 +585,7 @@ func (c *ChatCompletionClient) doStreamingRequest(ctx context.Context, apiKey st
return resp, nil
}

func (c *ChatCompletionClient) streamResponseToChannel(ctx context.Context, apiKey string, request openai.ChatCompletionRequest, dataChan chan<- string, errChan chan<- error, wg *sync.WaitGroup, opts []StreamOption) {
func (c *ChatCompletionClient) streamResponseToChannel(ctx context.Context, apiKey string, request CompletionRequest, dataChan chan<- string, errChan chan<- error, wg *sync.WaitGroup, opts []StreamOption) {
defer wg.Done()

resp, err := c.doStreamingRequest(ctx, apiKey, request, opts...)
Expand Down Expand Up @@ -640,7 +725,7 @@ func (c *ChatCompletionClient) handleStreamingToolCall(toolCall *openai.ToolCall
}
}

func (c *ChatCompletionClient) buildCompleteResponse(content string, reasoning string, functionCallAccumulator map[int]*functionCallAccumulator, toolCallAccumulator map[int]*toolCallAccumulator, model string, request openai.ChatCompletionRequest) openai.ChatCompletionResponse {
func (c *ChatCompletionClient) buildCompleteResponse(content string, reasoning string, functionCallAccumulator map[int]*functionCallAccumulator, toolCallAccumulator map[int]*toolCallAccumulator, model string, request CompletionRequest) openai.ChatCompletionResponse {
message := openai.ChatCompletionMessage{
Role: openai.ChatMessageRoleAssistant,
Content: content,
Expand Down
Loading
Loading