Skip to content
Open
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
176 changes: 113 additions & 63 deletions src/prometheus/ast.lua
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,44 @@ local astKindExpressionLookup = {

Ast.AstKind = AstKind;

-- Lua numbers use IEEE-754 doubles, so only integers through 2^53 - 1 are safe.
local MAX_SAFE_INT = 9007199254740991 -- 2^53 - 1
local MIN_SAFE_INT = -9007199254740991

local function isSafeInteger(n)
return type(n) == "number"
and n >= MIN_SAFE_INT
and n <= MAX_SAFE_INT
and (n % 1 == 0)
end

local function canCompareEqualityConstants(a, b)
if type(a) ~= type(b) then
return true
end
if type(a) == "number" then
return (a == a) and (b == b)
and a >= MIN_SAFE_INT and a <= MAX_SAFE_INT
and b >= MIN_SAFE_INT and b <= MAX_SAFE_INT
end
return true
end

local function canCompareOrderConstants(a, b)
if type(a) ~= type(b) then
return false
end
if type(a) == "number" then
return (a == a) and (b == b)
and a >= MIN_SAFE_INT and a <= MAX_SAFE_INT
and b >= MIN_SAFE_INT and b <= MAX_SAFE_INT
end
if type(a) == "string" then
return true
end
return false
end

function Ast.astKindExpressionToNumber(kind)
return astKindExpressionLookup[kind] or 100;
end
Expand Down Expand Up @@ -432,11 +470,12 @@ function Ast.NilExpression()
}
end

function Ast.NumberExpression(value)
function Ast.NumberExpression(value, raw)
return {
kind = AstKind.NumberExpression,
isConstant = true,
value = value,
raw = raw,
}
end

Expand All @@ -449,39 +488,39 @@ function Ast.StringExpression(value)
end

function Ast.OrExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value or rhs.value end);
if success then
return Ast.ConstantNode(val);
end
end
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value or rhs.value end);
if success then
return Ast.ConstantNode(val);
end
end

return {
kind = AstKind.OrExpression,
lhs = lhs,
rhs = rhs,
isConstant = false,
}
return {
kind = AstKind.OrExpression,
lhs = lhs,
rhs = rhs,
isConstant = false,
}
end

function Ast.AndExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value and rhs.value end);
if success then
return Ast.ConstantNode(val);
end
end
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value and rhs.value end);
if success then
return Ast.ConstantNode(val);
end
end

return {
kind = AstKind.AndExpression,
lhs = lhs,
rhs = rhs,
isConstant = false,
}
return {
kind = AstKind.AndExpression,
lhs = lhs,
rhs = rhs,
isConstant = false,
}
end

function Ast.LessThanExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
if(simplify and rhs.isConstant and lhs.isConstant and canCompareOrderConstants(lhs.value, rhs.value)) then
local success, val = pcall(function() return lhs.value < rhs.value end);
if success then
return Ast.ConstantNode(val);
Expand All @@ -497,7 +536,7 @@ function Ast.LessThanExpression(lhs, rhs, simplify)
end

function Ast.GreaterThanExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
if(simplify and rhs.isConstant and lhs.isConstant and canCompareOrderConstants(lhs.value, rhs.value)) then
local success, val = pcall(function() return lhs.value > rhs.value end);
if success then
return Ast.ConstantNode(val);
Expand All @@ -513,7 +552,7 @@ function Ast.GreaterThanExpression(lhs, rhs, simplify)
end

function Ast.LessThanOrEqualsExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
if(simplify and rhs.isConstant and lhs.isConstant and canCompareOrderConstants(lhs.value, rhs.value)) then
local success, val = pcall(function() return lhs.value <= rhs.value end);
if success then
return Ast.ConstantNode(val);
Expand All @@ -529,7 +568,7 @@ function Ast.LessThanOrEqualsExpression(lhs, rhs, simplify)
end

function Ast.GreaterThanOrEqualsExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
if(simplify and rhs.isConstant and lhs.isConstant and canCompareOrderConstants(lhs.value, rhs.value)) then
local success, val = pcall(function() return lhs.value >= rhs.value end);
if success then
return Ast.ConstantNode(val);
Expand All @@ -545,7 +584,7 @@ function Ast.GreaterThanOrEqualsExpression(lhs, rhs, simplify)
end

function Ast.NotEqualsExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
if(simplify and rhs.isConstant and lhs.isConstant and canCompareEqualityConstants(lhs.value, rhs.value)) then
local success, val = pcall(function() return lhs.value ~= rhs.value end);
if success then
return Ast.ConstantNode(val);
Expand All @@ -561,7 +600,7 @@ function Ast.NotEqualsExpression(lhs, rhs, simplify)
end

function Ast.EqualsExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
if(simplify and rhs.isConstant and lhs.isConstant and canCompareEqualityConstants(lhs.value, rhs.value)) then
local success, val = pcall(function() return lhs.value == rhs.value end);
if success then
return Ast.ConstantNode(val);
Expand All @@ -578,9 +617,8 @@ end

function Ast.StrCatExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value .. rhs.value end);
if success then
return Ast.ConstantNode(val);
if type(lhs.value) == "string" and type(rhs.value) == "string" then
return Ast.ConstantNode(lhs.value .. rhs.value);
end
end

Expand All @@ -594,9 +632,11 @@ end

function Ast.AddExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value + rhs.value end);
if success then
return Ast.ConstantNode(val);
if isSafeInteger(lhs.value) and isSafeInteger(rhs.value) then
local val = lhs.value + rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end

Expand All @@ -610,9 +650,11 @@ end

function Ast.SubExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value - rhs.value end);
if success then
return Ast.ConstantNode(val);
if isSafeInteger(lhs.value) and isSafeInteger(rhs.value) then
local val = lhs.value - rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end

Expand All @@ -626,9 +668,11 @@ end

function Ast.MulExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value * rhs.value end);
if success then
return Ast.ConstantNode(val);
if isSafeInteger(lhs.value) and isSafeInteger(rhs.value) then
local val = lhs.value * rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end

Expand All @@ -641,10 +685,14 @@ function Ast.MulExpression(lhs, rhs, simplify)
end

function Ast.DivExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant and rhs.value ~= 0) then
local success, val = pcall(function() return lhs.value / rhs.value end);
if success then
return Ast.ConstantNode(val);
if(simplify and rhs.isConstant and lhs.isConstant) then
if isSafeInteger(lhs.value) and isSafeInteger(rhs.value) and rhs.value ~= 0 then
if lhs.value % rhs.value == 0 then
local val = lhs.value / rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end
end

Expand All @@ -658,9 +706,11 @@ end

function Ast.ModExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value % rhs.value end);
if success then
return Ast.ConstantNode(val);
if isSafeInteger(lhs.value) and isSafeInteger(rhs.value) and rhs.value ~= 0 then
local val = lhs.value % rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end

Expand All @@ -674,10 +724,7 @@ end

function Ast.NotExpression(rhs, simplify)
if(simplify and rhs.isConstant) then
local success, val = pcall(function() return not rhs.value end);
if success then
return Ast.ConstantNode(val);
end
return Ast.ConstantNode(not rhs.value);
end

return {
Expand All @@ -689,9 +736,11 @@ end

function Ast.NegateExpression(rhs, simplify)
if(simplify and rhs.isConstant) then
local success, val = pcall(function() return -rhs.value end);
if success then
return Ast.ConstantNode(val);
if isSafeInteger(rhs.value) then
local val = -rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end

Expand All @@ -704,9 +753,8 @@ end

function Ast.LenExpression(rhs, simplify)
if(simplify and rhs.isConstant) then
local success, val = pcall(function() return #rhs.value end);
if success then
return Ast.ConstantNode(val);
if type(rhs.value) == "string" then
return Ast.ConstantNode(#rhs.value);
end
end

Expand All @@ -719,9 +767,11 @@ end

function Ast.PowExpression(lhs, rhs, simplify)
if(simplify and rhs.isConstant and lhs.isConstant) then
local success, val = pcall(function() return lhs.value ^ rhs.value end);
if success then
return Ast.ConstantNode(val);
if isSafeInteger(lhs.value) and isSafeInteger(rhs.value) and rhs.value >= 0 then
local val = lhs.value ^ rhs.value;
if isSafeInteger(val) then
return Ast.ConstantNode(val);
end
end
end

Expand Down
3 changes: 2 additions & 1 deletion src/prometheus/parser.lua
Original file line number Diff line number Diff line change
Expand Up @@ -908,7 +908,8 @@ function Parser:expressionLiteral(scope)

-- Number Literal
if(is(self, TokenKind.Number)) then
return Ast.NumberExpression(get(self).value);
local tk = get(self);
return Ast.NumberExpression(tk.value, tk.source);
end

-- True Literal
Expand Down
Loading