mirror of
https://github.com/RooVetGit/Roo-Code.git
synced 2026-08-28 05:27:24 +00:00
Merge pull request #863 from RooVetGit/temperature_control
Configure per-configuration temperature
This commit is contained in:
commit
ceec089745
19 changed files with 250 additions and 34 deletions
|
|
@ -1,7 +1,5 @@
|
|||
import { OpenAiHandler } from "../openai"
|
||||
import { ApiHandlerOptions } from "../../../shared/api"
|
||||
import { ApiStream } from "../../transform/stream"
|
||||
import OpenAI from "openai"
|
||||
import { Anthropic } from "@anthropic-ai/sdk"
|
||||
|
||||
// Mock OpenAI client
|
||||
|
|
|
|||
|
|
@ -10,6 +10,8 @@ import {
|
|||
import { ApiHandler, SingleCompletionHandler } from "../index"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
|
||||
const ANTHROPIC_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
|
||||
private options: ApiHandlerOptions
|
||||
private client: Anthropic
|
||||
|
|
@ -44,7 +46,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
|
|||
{
|
||||
model: modelId,
|
||||
max_tokens: this.getModel().info.maxTokens || 8192,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? ANTHROPIC_DEFAULT_TEMPERATURE,
|
||||
system: [{ text: systemPrompt, type: "text", cache_control: { type: "ephemeral" } }], // setting cache breakpoint for system prompt so new tasks can reuse it
|
||||
messages: messages.map((message, index) => {
|
||||
if (index === lastUserMsgIndex || index === secondLastMsgUserIndex) {
|
||||
|
|
@ -96,7 +98,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
|
|||
stream = (await this.client.messages.create({
|
||||
model: modelId,
|
||||
max_tokens: this.getModel().info.maxTokens || 8192,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? ANTHROPIC_DEFAULT_TEMPERATURE,
|
||||
system: [{ text: systemPrompt, type: "text" }],
|
||||
messages,
|
||||
// tools,
|
||||
|
|
@ -179,7 +181,7 @@ export class AnthropicHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const response = await this.client.messages.create({
|
||||
model: this.getModel().id,
|
||||
max_tokens: this.getModel().info.maxTokens || 8192,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? ANTHROPIC_DEFAULT_TEMPERATURE,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
stream: false,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -11,6 +11,8 @@ import { ApiHandlerOptions, BedrockModelId, ModelInfo, bedrockDefaultModelId, be
|
|||
import { ApiStream } from "../transform/stream"
|
||||
import { convertToBedrockConverseMessages, convertToAnthropicMessage } from "../transform/bedrock-converse-format"
|
||||
|
||||
const BEDROCK_DEFAULT_TEMPERATURE = 0.3
|
||||
|
||||
// Define types for stream events based on AWS SDK
|
||||
export interface StreamEvent {
|
||||
messageStart?: {
|
||||
|
|
@ -104,7 +106,7 @@ export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler {
|
|||
system: [{ text: systemPrompt }],
|
||||
inferenceConfig: {
|
||||
maxTokens: modelConfig.info.maxTokens || 5000,
|
||||
temperature: 0.3,
|
||||
temperature: this.options.modelTemperature ?? BEDROCK_DEFAULT_TEMPERATURE,
|
||||
topP: 0.1,
|
||||
...(this.options.awsUsePromptCache
|
||||
? {
|
||||
|
|
@ -262,7 +264,7 @@ export class AwsBedrockHandler implements ApiHandler, SingleCompletionHandler {
|
|||
]),
|
||||
inferenceConfig: {
|
||||
maxTokens: modelConfig.info.maxTokens || 5000,
|
||||
temperature: 0.3,
|
||||
temperature: this.options.modelTemperature ?? BEDROCK_DEFAULT_TEMPERATURE,
|
||||
topP: 0.1,
|
||||
},
|
||||
}
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ import { ApiHandlerOptions, geminiDefaultModelId, GeminiModelId, geminiModels, M
|
|||
import { convertAnthropicMessageToGemini } from "../transform/gemini-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
|
||||
const GEMINI_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class GeminiHandler implements ApiHandler, SingleCompletionHandler {
|
||||
private options: ApiHandlerOptions
|
||||
private client: GoogleGenerativeAI
|
||||
|
|
@ -23,7 +25,7 @@ export class GeminiHandler implements ApiHandler, SingleCompletionHandler {
|
|||
contents: messages.map(convertAnthropicMessageToGemini),
|
||||
generationConfig: {
|
||||
// maxOutputTokens: this.getModel().info.maxTokens,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? GEMINI_DEFAULT_TEMPERATURE,
|
||||
},
|
||||
})
|
||||
|
||||
|
|
@ -60,7 +62,7 @@ export class GeminiHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const result = await model.generateContent({
|
||||
contents: [{ role: "user", parts: [{ text: prompt }] }],
|
||||
generationConfig: {
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? GEMINI_DEFAULT_TEMPERATURE,
|
||||
},
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ import { ApiHandler, SingleCompletionHandler } from "../"
|
|||
import { ApiHandlerOptions, ModelInfo, glamaDefaultModelId, glamaDefaultModelInfo } from "../../shared/api"
|
||||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
import delay from "delay"
|
||||
|
||||
const GLAMA_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class GlamaHandler implements ApiHandler, SingleCompletionHandler {
|
||||
private options: ApiHandlerOptions
|
||||
|
|
@ -79,7 +80,7 @@ export class GlamaHandler implements ApiHandler, SingleCompletionHandler {
|
|||
}
|
||||
|
||||
if (this.supportsTemperature()) {
|
||||
requestOptions.temperature = 0
|
||||
requestOptions.temperature = this.options.modelTemperature ?? GLAMA_DEFAULT_TEMPERATURE
|
||||
}
|
||||
|
||||
const { data: completion, response } = await this.client.chat.completions
|
||||
|
|
@ -172,7 +173,7 @@ export class GlamaHandler implements ApiHandler, SingleCompletionHandler {
|
|||
}
|
||||
|
||||
if (this.supportsTemperature()) {
|
||||
requestOptions.temperature = 0
|
||||
requestOptions.temperature = this.options.modelTemperature ?? GLAMA_DEFAULT_TEMPERATURE
|
||||
}
|
||||
|
||||
if (this.getModel().id.startsWith("anthropic/")) {
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ import { ApiHandlerOptions, ModelInfo, openAiModelInfoSaneDefaults } from "../..
|
|||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
|
||||
const LMSTUDIO_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class LmStudioHandler implements ApiHandler, SingleCompletionHandler {
|
||||
private options: ApiHandlerOptions
|
||||
private client: OpenAI
|
||||
|
|
@ -27,7 +29,7 @@ export class LmStudioHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const stream = await this.client.chat.completions.create({
|
||||
model: this.getModel().id,
|
||||
messages: openAiMessages,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? LMSTUDIO_DEFAULT_TEMPERATURE,
|
||||
stream: true,
|
||||
})
|
||||
for await (const chunk of stream) {
|
||||
|
|
@ -59,7 +61,7 @@ export class LmStudioHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const response = await this.client.chat.completions.create({
|
||||
model: this.getModel().id,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? LMSTUDIO_DEFAULT_TEMPERATURE,
|
||||
stream: false,
|
||||
})
|
||||
return response.choices[0]?.message.content || ""
|
||||
|
|
|
|||
|
|
@ -14,6 +14,8 @@ import {
|
|||
import { convertToMistralMessages } from "../transform/mistral-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
|
||||
const MISTRAL_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class MistralHandler implements ApiHandler {
|
||||
private options: ApiHandlerOptions
|
||||
private client: Mistral
|
||||
|
|
@ -30,7 +32,7 @@ export class MistralHandler implements ApiHandler {
|
|||
const stream = await this.client.chat.stream({
|
||||
model: this.getModel().id,
|
||||
// max_completion_tokens: this.getModel().info.maxTokens,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? MISTRAL_DEFAULT_TEMPERATURE,
|
||||
messages: [{ role: "system", content: systemPrompt }, ...convertToMistralMessages(messages)],
|
||||
stream: true,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -5,6 +5,9 @@ import { ApiHandlerOptions, ModelInfo, openAiModelInfoSaneDefaults } from "../..
|
|||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
import { convertToR1Format } from "../transform/r1-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
import { DEEP_SEEK_DEFAULT_TEMPERATURE } from "./openai"
|
||||
|
||||
const OLLAMA_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
|
||||
private options: ApiHandlerOptions
|
||||
|
|
@ -20,7 +23,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
|
|||
|
||||
async *createMessage(systemPrompt: string, messages: Anthropic.Messages.MessageParam[]): ApiStream {
|
||||
const modelId = this.getModel().id
|
||||
const useR1Format = modelId.toLowerCase().includes('deepseek-r1')
|
||||
const useR1Format = modelId.toLowerCase().includes("deepseek-r1")
|
||||
const openAiMessages: OpenAI.Chat.ChatCompletionMessageParam[] = [
|
||||
{ role: "system", content: systemPrompt },
|
||||
...(useR1Format ? convertToR1Format(messages) : convertToOpenAiMessages(messages)),
|
||||
|
|
@ -29,7 +32,7 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const stream = await this.client.chat.completions.create({
|
||||
model: this.getModel().id,
|
||||
messages: openAiMessages,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? OLLAMA_DEFAULT_TEMPERATURE,
|
||||
stream: true,
|
||||
})
|
||||
for await (const chunk of stream) {
|
||||
|
|
@ -53,11 +56,15 @@ export class OllamaHandler implements ApiHandler, SingleCompletionHandler {
|
|||
async completePrompt(prompt: string): Promise<string> {
|
||||
try {
|
||||
const modelId = this.getModel().id
|
||||
const useR1Format = modelId.toLowerCase().includes('deepseek-r1')
|
||||
const useR1Format = modelId.toLowerCase().includes("deepseek-r1")
|
||||
const response = await this.client.chat.completions.create({
|
||||
model: this.getModel().id,
|
||||
messages: useR1Format ? convertToR1Format([{ role: "user", content: prompt }]) : [{ role: "user", content: prompt }],
|
||||
temperature: 0,
|
||||
messages: useR1Format
|
||||
? convertToR1Format([{ role: "user", content: prompt }])
|
||||
: [{ role: "user", content: prompt }],
|
||||
temperature:
|
||||
this.options.modelTemperature ??
|
||||
(useR1Format ? DEEP_SEEK_DEFAULT_TEMPERATURE : OLLAMA_DEFAULT_TEMPERATURE),
|
||||
stream: false,
|
||||
})
|
||||
return response.choices[0]?.message.content || ""
|
||||
|
|
|
|||
|
|
@ -11,6 +11,8 @@ import {
|
|||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
|
||||
const OPENAI_NATIVE_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler {
|
||||
private options: ApiHandlerOptions
|
||||
private client: OpenAI
|
||||
|
|
@ -88,7 +90,7 @@ export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler
|
|||
): ApiStream {
|
||||
const stream = await this.client.chat.completions.create({
|
||||
model: modelId,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? OPENAI_NATIVE_DEFAULT_TEMPERATURE,
|
||||
messages: [{ role: "system", content: systemPrompt }, ...convertToOpenAiMessages(messages)],
|
||||
stream: true,
|
||||
stream_options: { include_usage: true },
|
||||
|
|
@ -189,7 +191,7 @@ export class OpenAiNativeHandler implements ApiHandler, SingleCompletionHandler
|
|||
return {
|
||||
model: modelId,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? OPENAI_NATIVE_DEFAULT_TEMPERATURE,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -13,6 +13,9 @@ import { convertToR1Format } from "../transform/r1-format"
|
|||
import { convertToSimpleMessages } from "../transform/simple-format"
|
||||
import { ApiStream } from "../transform/stream"
|
||||
|
||||
export const DEEP_SEEK_DEFAULT_TEMPERATURE = 0.6
|
||||
const OPENAI_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
export class OpenAiHandler implements ApiHandler, SingleCompletionHandler {
|
||||
protected options: ApiHandlerOptions
|
||||
private client: OpenAI
|
||||
|
|
@ -70,7 +73,9 @@ export class OpenAiHandler implements ApiHandler, SingleCompletionHandler {
|
|||
|
||||
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsStreaming = {
|
||||
model: modelId,
|
||||
temperature: 0,
|
||||
temperature:
|
||||
this.options.modelTemperature ??
|
||||
(deepseekReasoner ? DEEP_SEEK_DEFAULT_TEMPERATURE : OPENAI_DEFAULT_TEMPERATURE),
|
||||
messages: convertedMessages,
|
||||
stream: true as const,
|
||||
stream_options: { include_usage: true },
|
||||
|
|
|
|||
|
|
@ -6,6 +6,9 @@ import { ApiHandlerOptions, ModelInfo, openRouterDefaultModelId, openRouterDefau
|
|||
import { convertToOpenAiMessages } from "../transform/openai-format"
|
||||
import { ApiStreamChunk, ApiStreamUsageChunk } from "../transform/stream"
|
||||
import delay from "delay"
|
||||
import { DEEP_SEEK_DEFAULT_TEMPERATURE } from "./openai"
|
||||
|
||||
const OPENROUTER_DEFAULT_TEMPERATURE = 0
|
||||
|
||||
// Add custom interface for OpenRouter params
|
||||
type OpenRouterChatCompletionParams = OpenAI.Chat.ChatCompletionCreateParams & {
|
||||
|
|
@ -115,7 +118,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
|
|||
break
|
||||
}
|
||||
|
||||
let temperature = 0
|
||||
let defaultTemperature = OPENROUTER_DEFAULT_TEMPERATURE
|
||||
let topP: number | undefined = undefined
|
||||
|
||||
// Handle models based on deepseek-r1
|
||||
|
|
@ -124,9 +127,8 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
|
|||
this.getModel().id === "perplexity/sonar-reasoning"
|
||||
) {
|
||||
// Recommended temperature for DeepSeek reasoning models
|
||||
temperature = 0.6
|
||||
// DeepSeek highly recommends using user instead of system
|
||||
// role
|
||||
defaultTemperature = DEEP_SEEK_DEFAULT_TEMPERATURE
|
||||
// DeepSeek highly recommends using user instead of system role
|
||||
openAiMessages = convertToR1Format([{ role: "user", content: systemPrompt }, ...messages])
|
||||
// Some provider support topP and 0.95 is value that Deepseek used in their benchmarks
|
||||
topP = 0.95
|
||||
|
|
@ -137,7 +139,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const stream = await this.client.chat.completions.create({
|
||||
model: this.getModel().id,
|
||||
max_tokens: maxTokens,
|
||||
temperature: temperature,
|
||||
temperature: this.options.modelTemperature ?? defaultTemperature,
|
||||
top_p: topP,
|
||||
messages: openAiMessages,
|
||||
stream: true,
|
||||
|
|
@ -224,7 +226,7 @@ export class OpenRouterHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const response = await this.client.chat.completions.create({
|
||||
model: this.getModel().id,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? OPENROUTER_DEFAULT_TEMPERATURE,
|
||||
stream: false,
|
||||
})
|
||||
|
||||
|
|
|
|||
|
|
@ -79,7 +79,7 @@ export class UnboundHandler implements ApiHandler, SingleCompletionHandler {
|
|||
{
|
||||
model: this.getModel().id.split("/")[1],
|
||||
max_tokens: maxTokens,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
messages: openAiMessages,
|
||||
stream: true,
|
||||
},
|
||||
|
|
@ -146,7 +146,7 @@ export class UnboundHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const requestOptions: OpenAI.Chat.Completions.ChatCompletionCreateParamsNonStreaming = {
|
||||
model: this.getModel().id.split("/")[1],
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
}
|
||||
|
||||
if (this.getModel().id.startsWith("anthropic/")) {
|
||||
|
|
|
|||
|
|
@ -22,7 +22,7 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const stream = await this.client.messages.create({
|
||||
model: this.getModel().id,
|
||||
max_tokens: this.getModel().info.maxTokens || 8192,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
system: systemPrompt,
|
||||
messages,
|
||||
stream: true,
|
||||
|
|
@ -89,7 +89,7 @@ export class VertexHandler implements ApiHandler, SingleCompletionHandler {
|
|||
const response = await this.client.messages.create({
|
||||
model: this.getModel().id,
|
||||
max_tokens: this.getModel().info.maxTokens || 8192,
|
||||
temperature: 0,
|
||||
temperature: this.options.modelTemperature ?? 0,
|
||||
messages: [{ role: "user", content: prompt }],
|
||||
stream: false,
|
||||
})
|
||||
|
|
|
|||
|
|
@ -123,6 +123,7 @@ type GlobalStateKey =
|
|||
| "customModes" // Array of custom modes
|
||||
| "unboundModelId"
|
||||
| "unboundModelInfo"
|
||||
| "modelTemperature"
|
||||
|
||||
export const GlobalFileNames = {
|
||||
apiConversationHistory: "api_conversation_history.json",
|
||||
|
|
@ -1587,6 +1588,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
|
|||
unboundApiKey,
|
||||
unboundModelId,
|
||||
unboundModelInfo,
|
||||
modelTemperature,
|
||||
} = apiConfiguration
|
||||
await this.updateGlobalState("apiProvider", apiProvider)
|
||||
await this.updateGlobalState("apiModelId", apiModelId)
|
||||
|
|
@ -1628,6 +1630,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
|
|||
await this.storeSecret("unboundApiKey", unboundApiKey)
|
||||
await this.updateGlobalState("unboundModelId", unboundModelId)
|
||||
await this.updateGlobalState("unboundModelInfo", unboundModelInfo)
|
||||
await this.updateGlobalState("modelTemperature", modelTemperature)
|
||||
if (this.cline) {
|
||||
this.cline.api = buildApiHandler(apiConfiguration)
|
||||
}
|
||||
|
|
@ -2388,6 +2391,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
|
|||
unboundApiKey,
|
||||
unboundModelId,
|
||||
unboundModelInfo,
|
||||
modelTemperature,
|
||||
] = await Promise.all([
|
||||
this.getGlobalState("apiProvider") as Promise<ApiProvider | undefined>,
|
||||
this.getGlobalState("apiModelId") as Promise<string | undefined>,
|
||||
|
|
@ -2464,6 +2468,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
|
|||
this.getSecret("unboundApiKey") as Promise<string | undefined>,
|
||||
this.getGlobalState("unboundModelId") as Promise<string | undefined>,
|
||||
this.getGlobalState("unboundModelInfo") as Promise<ModelInfo | undefined>,
|
||||
this.getGlobalState("modelTemperature") as Promise<number | undefined>,
|
||||
])
|
||||
|
||||
let apiProvider: ApiProvider
|
||||
|
|
@ -2522,6 +2527,7 @@ export class ClineProvider implements vscode.WebviewViewProvider {
|
|||
unboundApiKey,
|
||||
unboundModelId,
|
||||
unboundModelInfo,
|
||||
modelTemperature,
|
||||
},
|
||||
lastShownAnnouncementId,
|
||||
customInstructions,
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ export interface ApiHandlerOptions {
|
|||
unboundApiKey?: string
|
||||
unboundModelId?: string
|
||||
unboundModelInfo?: ModelInfo
|
||||
modelTemperature?: number
|
||||
}
|
||||
|
||||
export type ApiConfiguration = ApiHandlerOptions & {
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ import { memo, useCallback, useEffect, useMemo, useState } from "react"
|
|||
import { useEvent, useInterval } from "react-use"
|
||||
import { Checkbox, Dropdown, Pane, type DropdownOption } from "vscrui"
|
||||
import { VSCodeLink, VSCodeRadio, VSCodeRadioGroup, VSCodeTextField } from "@vscode/webview-ui-toolkit/react"
|
||||
import { TemperatureControl } from "./TemperatureControl"
|
||||
import * as vscodemodels from "vscode"
|
||||
|
||||
import {
|
||||
|
|
@ -1363,6 +1364,18 @@ const ApiOptions = ({ apiErrorMessage, modelIdErrorMessage }: ApiOptionsProps) =
|
|||
</>
|
||||
)}
|
||||
|
||||
<div style={{ marginTop: "10px" }}>
|
||||
<TemperatureControl
|
||||
value={apiConfiguration?.modelTemperature}
|
||||
onChange={(value) => {
|
||||
handleInputChange("modelTemperature")({
|
||||
target: { value },
|
||||
})
|
||||
}}
|
||||
maxValue={2}
|
||||
/>
|
||||
</div>
|
||||
|
||||
{modelIdErrorMessage && (
|
||||
<p
|
||||
style={{
|
||||
|
|
|
|||
80
webview-ui/src/components/settings/TemperatureControl.tsx
Normal file
80
webview-ui/src/components/settings/TemperatureControl.tsx
Normal file
|
|
@ -0,0 +1,80 @@
|
|||
import { VSCodeCheckbox } from "@vscode/webview-ui-toolkit/react"
|
||||
import { useEffect, useState } from "react"
|
||||
|
||||
interface TemperatureControlProps {
|
||||
value: number | undefined
|
||||
onChange: (value: number | undefined) => void
|
||||
maxValue?: number // Some providers like OpenAI use 0-2 range
|
||||
}
|
||||
|
||||
export const TemperatureControl = ({ value, onChange, maxValue = 1 }: TemperatureControlProps) => {
|
||||
const [isCustomTemperature, setIsCustomTemperature] = useState(value !== undefined)
|
||||
const [inputValue, setInputValue] = useState(value?.toString() ?? "0")
|
||||
|
||||
// Sync internal state with prop changes when switching profiles
|
||||
useEffect(() => {
|
||||
const hasCustomTemperature = value !== undefined
|
||||
setIsCustomTemperature(hasCustomTemperature)
|
||||
setInputValue(value?.toString() ?? "0")
|
||||
}, [value])
|
||||
|
||||
return (
|
||||
<div>
|
||||
<VSCodeCheckbox
|
||||
checked={isCustomTemperature}
|
||||
onChange={(e: any) => {
|
||||
const isChecked = e.target.checked
|
||||
setIsCustomTemperature(isChecked)
|
||||
if (!isChecked) {
|
||||
onChange(undefined) // Unset the temperature
|
||||
} else if (value !== undefined) {
|
||||
onChange(value) // Use the value from apiConfiguration, if set
|
||||
}
|
||||
}}>
|
||||
<span style={{ fontWeight: "500" }}>Use custom temperature</span>
|
||||
</VSCodeCheckbox>
|
||||
|
||||
<p style={{ fontSize: "12px", marginTop: "5px", color: "var(--vscode-descriptionForeground)" }}>
|
||||
Controls randomness in the model's responses.
|
||||
</p>
|
||||
|
||||
{isCustomTemperature && (
|
||||
<div
|
||||
style={{
|
||||
marginTop: 5,
|
||||
marginBottom: 10,
|
||||
paddingLeft: 10,
|
||||
borderLeft: "2px solid var(--vscode-button-background)",
|
||||
}}>
|
||||
<div style={{ display: "flex", alignItems: "center", gap: "5px" }}>
|
||||
<input
|
||||
aria-label="Temperature control text input"
|
||||
type="text"
|
||||
value={inputValue}
|
||||
onChange={(e) => setInputValue(e.target.value)}
|
||||
onBlur={(e) => {
|
||||
const newValue = parseFloat(e.target.value)
|
||||
if (!isNaN(newValue) && newValue >= 0 && newValue <= maxValue) {
|
||||
onChange(newValue)
|
||||
setInputValue(newValue.toString())
|
||||
} else {
|
||||
setInputValue(value?.toString() ?? "0") // Reset to last valid value
|
||||
}
|
||||
}}
|
||||
style={{
|
||||
width: "60px",
|
||||
padding: "4px 8px",
|
||||
border: "1px solid var(--vscode-input-border)",
|
||||
background: "var(--vscode-input-background)",
|
||||
color: "var(--vscode-input-foreground)",
|
||||
}}
|
||||
/>
|
||||
</div>
|
||||
<p style={{ fontSize: "12px", marginTop: "8px", color: "var(--vscode-descriptionForeground)" }}>
|
||||
Higher values make output more random, lower values make it more deterministic.
|
||||
</p>
|
||||
</div>
|
||||
)}
|
||||
</div>
|
||||
)
|
||||
}
|
||||
|
|
@ -0,0 +1,86 @@
|
|||
import { render, screen, fireEvent } from "@testing-library/react"
|
||||
import { TemperatureControl } from "../TemperatureControl"
|
||||
|
||||
describe("TemperatureControl", () => {
|
||||
it("renders with default temperature disabled", () => {
|
||||
const onChange = jest.fn()
|
||||
render(<TemperatureControl value={undefined} onChange={onChange} />)
|
||||
|
||||
const checkbox = screen.getByRole("checkbox")
|
||||
expect(checkbox).not.toBeChecked()
|
||||
expect(screen.queryByRole("textbox")).not.toBeInTheDocument()
|
||||
})
|
||||
|
||||
it("renders with custom temperature enabled", () => {
|
||||
const onChange = jest.fn()
|
||||
render(<TemperatureControl value={0.7} onChange={onChange} />)
|
||||
|
||||
const checkbox = screen.getByRole("checkbox")
|
||||
expect(checkbox).toBeChecked()
|
||||
|
||||
const input = screen.getByRole("textbox")
|
||||
expect(input).toBeInTheDocument()
|
||||
expect(input).toHaveValue("0.7")
|
||||
})
|
||||
|
||||
it("updates when checkbox is toggled", () => {
|
||||
const onChange = jest.fn()
|
||||
render(<TemperatureControl value={0.7} onChange={onChange} />)
|
||||
|
||||
const checkbox = screen.getByRole("checkbox")
|
||||
|
||||
// Uncheck - should clear temperature
|
||||
fireEvent.click(checkbox)
|
||||
expect(onChange).toHaveBeenCalledWith(undefined)
|
||||
|
||||
// Check - should restore previous temperature
|
||||
fireEvent.click(checkbox)
|
||||
expect(onChange).toHaveBeenCalledWith(0.7)
|
||||
})
|
||||
|
||||
it("updates temperature when input loses focus", () => {
|
||||
const onChange = jest.fn()
|
||||
render(<TemperatureControl value={0.7} onChange={onChange} />)
|
||||
|
||||
const input = screen.getByRole("textbox")
|
||||
fireEvent.change(input, { target: { value: "0.8" } })
|
||||
fireEvent.blur(input)
|
||||
|
||||
expect(onChange).toHaveBeenCalledWith(0.8)
|
||||
})
|
||||
|
||||
it("respects maxValue prop", () => {
|
||||
const onChange = jest.fn()
|
||||
render(<TemperatureControl value={1.5} onChange={onChange} maxValue={2} />)
|
||||
|
||||
const input = screen.getByRole("textbox")
|
||||
|
||||
// Valid value within max
|
||||
fireEvent.change(input, { target: { value: "1.8" } })
|
||||
fireEvent.blur(input)
|
||||
expect(onChange).toHaveBeenCalledWith(1.8)
|
||||
|
||||
// Invalid value above max
|
||||
fireEvent.change(input, { target: { value: "2.5" } })
|
||||
fireEvent.blur(input)
|
||||
expect(input).toHaveValue("1.5") // Should revert to original value
|
||||
expect(onChange).toHaveBeenCalledTimes(1) // Should not call onChange for invalid value
|
||||
})
|
||||
|
||||
it("syncs checkbox state when value prop changes", () => {
|
||||
const onChange = jest.fn()
|
||||
const { rerender } = render(<TemperatureControl value={0.7} onChange={onChange} />)
|
||||
|
||||
// Initially checked
|
||||
const checkbox = screen.getByRole("checkbox")
|
||||
expect(checkbox).toBeChecked()
|
||||
|
||||
// Update to undefined
|
||||
rerender(<TemperatureControl value={undefined} onChange={onChange} />)
|
||||
expect(checkbox).not.toBeChecked()
|
||||
|
||||
// Update back to a value
|
||||
rerender(<TemperatureControl value={0.5} onChange={onChange} />)
|
||||
expect(checkbox).toBeChecked()
|
||||
})
|
||||
})
|
||||
|
|
@ -1,4 +1,9 @@
|
|||
import { ApiConfiguration, glamaDefaultModelId, openRouterDefaultModelId, unboundDefaultModelId } from "../../../src/shared/api"
|
||||
import {
|
||||
ApiConfiguration,
|
||||
glamaDefaultModelId,
|
||||
openRouterDefaultModelId,
|
||||
unboundDefaultModelId,
|
||||
} from "../../../src/shared/api"
|
||||
import { ModelInfo } from "../../../src/shared/api"
|
||||
export function validateApiConfiguration(apiConfiguration?: ApiConfiguration): string | undefined {
|
||||
if (apiConfiguration) {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue