implement passing two action input properties to cover all model scenarios
This commit is contained in:
@@ -109,6 +109,7 @@ describe('helpers.ts - inference request building', () => {
|
|||||||
undefined,
|
undefined,
|
||||||
undefined,
|
undefined,
|
||||||
100,
|
100,
|
||||||
|
undefined,
|
||||||
'https://api.test.com',
|
'https://api.test.com',
|
||||||
'test-token',
|
'test-token',
|
||||||
)
|
)
|
||||||
@@ -122,6 +123,7 @@ describe('helpers.ts - inference request building', () => {
|
|||||||
temperature: undefined,
|
temperature: undefined,
|
||||||
topP: undefined,
|
topP: undefined,
|
||||||
maxTokens: 100,
|
maxTokens: 100,
|
||||||
|
maxCompletionTokens: undefined,
|
||||||
endpoint: 'https://api.test.com',
|
endpoint: 'https://api.test.com',
|
||||||
token: 'test-token',
|
token: 'test-token',
|
||||||
responseFormat: {
|
responseFormat: {
|
||||||
@@ -143,6 +145,7 @@ describe('helpers.ts - inference request building', () => {
|
|||||||
undefined,
|
undefined,
|
||||||
undefined,
|
undefined,
|
||||||
100,
|
100,
|
||||||
|
undefined,
|
||||||
'https://api.test.com',
|
'https://api.test.com',
|
||||||
'test-token',
|
'test-token',
|
||||||
)
|
)
|
||||||
@@ -156,6 +159,7 @@ describe('helpers.ts - inference request building', () => {
|
|||||||
temperature: undefined,
|
temperature: undefined,
|
||||||
topP: undefined,
|
topP: undefined,
|
||||||
maxTokens: 100,
|
maxTokens: 100,
|
||||||
|
maxCompletionTokens: undefined,
|
||||||
endpoint: 'https://api.test.com',
|
endpoint: 'https://api.test.com',
|
||||||
token: 'test-token',
|
token: 'test-token',
|
||||||
responseFormat: undefined,
|
responseFormat: undefined,
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ describe('inference.ts', () => {
|
|||||||
{role: 'user' as const, content: 'Hello, AI!'},
|
{role: 'user' as const, content: 'Hello, AI!'},
|
||||||
],
|
],
|
||||||
modelName: 'gpt-4',
|
modelName: 'gpt-4',
|
||||||
maxTokens: 100,
|
maxCompletionTokens: 100,
|
||||||
endpoint: 'https://api.test.com',
|
endpoint: 'https://api.test.com',
|
||||||
token: 'test-token',
|
token: 'test-token',
|
||||||
}
|
}
|
||||||
@@ -633,4 +633,35 @@ describe('inference.ts', () => {
|
|||||||
expect(result).toBe('{"immediate": "result"}')
|
expect(result).toBe('{"immediate": "result"}')
|
||||||
})
|
})
|
||||||
})
|
})
|
||||||
|
|
||||||
|
describe('token param routing', () => {
|
||||||
|
it('sends max_tokens when only maxTokens is set', async () => {
|
||||||
|
const requestWithMaxTokens = {
|
||||||
|
...mockRequest,
|
||||||
|
maxCompletionTokens: undefined,
|
||||||
|
maxTokens: 100,
|
||||||
|
}
|
||||||
|
|
||||||
|
const mockResponse = {
|
||||||
|
choices: [
|
||||||
|
{
|
||||||
|
message: {
|
||||||
|
content: 'Direct max_tokens response',
|
||||||
|
},
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
|
||||||
|
mockCreate.mockResolvedValueOnce(mockResponse)
|
||||||
|
|
||||||
|
const result = await simpleInference(requestWithMaxTokens)
|
||||||
|
|
||||||
|
expect(result).toBe('Direct max_tokens response')
|
||||||
|
expect(mockCreate).toHaveBeenCalledTimes(1)
|
||||||
|
|
||||||
|
// Should have sent max_tokens directly
|
||||||
|
expect(mockCreate.mock.calls[0][0]).toHaveProperty('max_tokens', 100)
|
||||||
|
expect(mockCreate.mock.calls[0][0]).not.toHaveProperty('max_completion_tokens')
|
||||||
|
})
|
||||||
|
})
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -168,6 +168,7 @@ describe('main.ts', () => {
|
|||||||
],
|
],
|
||||||
modelName: 'gpt-4',
|
modelName: 'gpt-4',
|
||||||
maxTokens: 100,
|
maxTokens: 100,
|
||||||
|
maxCompletionTokens: undefined,
|
||||||
endpoint: 'https://api.test.com',
|
endpoint: 'https://api.test.com',
|
||||||
token: 'fake-token',
|
token: 'fake-token',
|
||||||
responseFormat: undefined,
|
responseFormat: undefined,
|
||||||
@@ -259,6 +260,7 @@ describe('main.ts', () => {
|
|||||||
],
|
],
|
||||||
modelName: 'gpt-4',
|
modelName: 'gpt-4',
|
||||||
maxTokens: 100,
|
maxTokens: 100,
|
||||||
|
maxCompletionTokens: undefined,
|
||||||
endpoint: 'https://api.test.com',
|
endpoint: 'https://api.test.com',
|
||||||
token: 'fake-token',
|
token: 'fake-token',
|
||||||
responseFormat: undefined,
|
responseFormat: undefined,
|
||||||
|
|||||||
+5
-1
@@ -43,9 +43,13 @@ inputs:
|
|||||||
required: false
|
required: false
|
||||||
default: ''
|
default: ''
|
||||||
max-tokens:
|
max-tokens:
|
||||||
description: The maximum number of tokens to generate
|
description: Maximum tokens to generate (deprecated)
|
||||||
required: false
|
required: false
|
||||||
default: '200'
|
default: '200'
|
||||||
|
max-completion-tokens:
|
||||||
|
description: Maximum tokens to generate
|
||||||
|
required: false
|
||||||
|
default: ''
|
||||||
temperature:
|
temperature:
|
||||||
description: The sampling temperature to use (0-1)
|
description: The sampling temperature to use (0-1)
|
||||||
required: false
|
required: false
|
||||||
|
|||||||
+4
-2
@@ -162,7 +162,8 @@ export function buildInferenceRequest(
|
|||||||
modelName: string,
|
modelName: string,
|
||||||
temperature: number | undefined,
|
temperature: number | undefined,
|
||||||
topP: number | undefined,
|
topP: number | undefined,
|
||||||
maxTokens: number,
|
maxTokens: number | undefined, // Deprecated
|
||||||
|
maxCompletionTokens: number | undefined,
|
||||||
endpoint: string,
|
endpoint: string,
|
||||||
token: string,
|
token: string,
|
||||||
customHeaders?: Record<string, string>,
|
customHeaders?: Record<string, string>,
|
||||||
@@ -175,7 +176,8 @@ export function buildInferenceRequest(
|
|||||||
modelName,
|
modelName,
|
||||||
temperature,
|
temperature,
|
||||||
topP,
|
topP,
|
||||||
maxTokens,
|
maxTokens, // Deprecated
|
||||||
|
maxCompletionTokens,
|
||||||
endpoint,
|
endpoint,
|
||||||
token,
|
token,
|
||||||
responseFormat,
|
responseFormat,
|
||||||
|
|||||||
+19
-6
@@ -12,7 +12,8 @@ interface ChatMessage {
|
|||||||
export interface InferenceRequest {
|
export interface InferenceRequest {
|
||||||
messages: Array<{role: 'system' | 'user' | 'assistant' | 'tool'; content: string}>
|
messages: Array<{role: 'system' | 'user' | 'assistant' | 'tool'; content: string}>
|
||||||
modelName: string
|
modelName: string
|
||||||
maxTokens: number
|
maxTokens?: number // Deprecated
|
||||||
|
maxCompletionTokens?: number
|
||||||
endpoint: string
|
endpoint: string
|
||||||
token: string
|
token: string
|
||||||
temperature?: number
|
temperature?: number
|
||||||
@@ -33,6 +34,20 @@ export interface InferenceResponse {
|
|||||||
}>
|
}>
|
||||||
}
|
}
|
||||||
|
|
||||||
|
/**
|
||||||
|
* Build the token limit params for a chat completion request.
|
||||||
|
* Only one of max_tokens or max_completion_tokens will be set.
|
||||||
|
*/
|
||||||
|
function buildMaxTokensParam(request: InferenceRequest): {max_tokens?: number; max_completion_tokens?: number} {
|
||||||
|
if (request.maxCompletionTokens != null) {
|
||||||
|
return {max_completion_tokens: request.maxCompletionTokens}
|
||||||
|
}
|
||||||
|
if (request.maxTokens != null) {
|
||||||
|
return {max_tokens: request.maxTokens}
|
||||||
|
}
|
||||||
|
return {}
|
||||||
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Simple one-shot inference without tools
|
* Simple one-shot inference without tools
|
||||||
*/
|
*/
|
||||||
@@ -47,10 +62,10 @@ export async function simpleInference(request: InferenceRequest): Promise<string
|
|||||||
|
|
||||||
const chatCompletionRequest: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
const chatCompletionRequest: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
||||||
messages: request.messages as OpenAI.Chat.Completions.ChatCompletionMessageParam[],
|
messages: request.messages as OpenAI.Chat.Completions.ChatCompletionMessageParam[],
|
||||||
max_completion_tokens: request.maxTokens,
|
|
||||||
model: request.modelName,
|
model: request.modelName,
|
||||||
temperature: request.temperature,
|
temperature: request.temperature,
|
||||||
top_p: request.topP,
|
top_p: request.topP,
|
||||||
|
...buildMaxTokensParam(request), // Note: solution around models using different underlying max tokens properties
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add response format if specified
|
// Add response format if specified
|
||||||
@@ -95,10 +110,10 @@ export async function mcpInference(
|
|||||||
|
|
||||||
const chatCompletionRequest: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
const chatCompletionRequest: OpenAI.Chat.Completions.ChatCompletionCreateParams = {
|
||||||
messages: messages as OpenAI.Chat.Completions.ChatCompletionMessageParam[],
|
messages: messages as OpenAI.Chat.Completions.ChatCompletionMessageParam[],
|
||||||
max_completion_tokens: request.maxTokens,
|
|
||||||
model: request.modelName,
|
model: request.modelName,
|
||||||
temperature: request.temperature,
|
temperature: request.temperature,
|
||||||
top_p: request.topP,
|
top_p: request.topP,
|
||||||
|
...buildMaxTokensParam(request),
|
||||||
}
|
}
|
||||||
|
|
||||||
// Add response format if specified (only on final iteration to avoid conflicts with tool calls)
|
// Add response format if specified (only on final iteration to avoid conflicts with tool calls)
|
||||||
@@ -162,9 +177,7 @@ export async function mcpInference(
|
|||||||
}
|
}
|
||||||
|
|
||||||
/**
|
/**
|
||||||
* Wrapper around OpenAI chat.completions.create with defensive handling for cases where
|
* Wrapper around OpenAI chat.completions.create with response validation.
|
||||||
* the SDK returns a raw string (e.g., unexpected content-type or streaming body) instead of
|
|
||||||
* a parsed object. Ensures an object with a 'choices' array is returned or throws a descriptive error.
|
|
||||||
*/
|
*/
|
||||||
async function chatCompletion(
|
async function chatCompletion(
|
||||||
client: OpenAI,
|
client: OpenAI,
|
||||||
|
|||||||
+8
-4
@@ -48,11 +48,14 @@ export async function run(): Promise<void> {
|
|||||||
|
|
||||||
// Get common parameters
|
// Get common parameters
|
||||||
const modelName = promptConfig?.model || core.getInput('model')
|
const modelName = promptConfig?.model || core.getInput('model')
|
||||||
let maxTokens = promptConfig?.modelParameters?.maxTokens ?? core.getInput('max-tokens')
|
|
||||||
|
|
||||||
if (typeof maxTokens === 'string') {
|
// Parse token limit inputs
|
||||||
maxTokens = parseInt(maxTokens, 10)
|
const maxCompletionTokensInput =
|
||||||
}
|
promptConfig?.modelParameters?.maxCompletionTokens ?? core.getInput('max-completion-tokens')
|
||||||
|
const maxCompletionTokens = maxCompletionTokensInput ? Number(maxCompletionTokensInput) : undefined
|
||||||
|
|
||||||
|
const maxTokensInput = promptConfig?.modelParameters?.maxTokens ?? core.getInput('max-tokens')
|
||||||
|
const maxTokens = maxCompletionTokens != null ? undefined : maxTokensInput ? Number(maxTokensInput) : undefined
|
||||||
|
|
||||||
const token = process.env['GITHUB_TOKEN'] || core.getInput('token')
|
const token = process.env['GITHUB_TOKEN'] || core.getInput('token')
|
||||||
if (token === undefined) {
|
if (token === undefined) {
|
||||||
@@ -85,6 +88,7 @@ export async function run(): Promise<void> {
|
|||||||
temperature,
|
temperature,
|
||||||
topP,
|
topP,
|
||||||
maxTokens,
|
maxTokens,
|
||||||
|
maxCompletionTokens,
|
||||||
endpoint,
|
endpoint,
|
||||||
token,
|
token,
|
||||||
customHeaders,
|
customHeaders,
|
||||||
|
|||||||
+2
-1
@@ -8,7 +8,8 @@ export interface PromptMessage {
|
|||||||
}
|
}
|
||||||
|
|
||||||
export interface ModelParameters {
|
export interface ModelParameters {
|
||||||
maxTokens?: number
|
maxTokens?: number // Deprecated
|
||||||
|
maxCompletionTokens?: number
|
||||||
temperature?: number
|
temperature?: number
|
||||||
topP?: number
|
topP?: number
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user