Skip to content
Merged
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
86 changes: 60 additions & 26 deletions VM/src/lstrlib.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1084,44 +1084,78 @@ static int str_split(lua_State* L)
size_t needleLen;
const char* needle = luaL_optlstring(L, 2, ",", &needleLen);

const char* begin = haystack;
const char* end = haystack + haystackLen;
const char* spanStart = begin;
const char* spanStart = haystack;
int numMatches = 0;

lua_createtable(L, 0, 0);

// Use of memchr/memcmp here instead of strchr/strncmp is so that we allow
// embedded nulls to be used in either of the haystack or the needle
// strings. Most Lua string APIs allow embedded nulls, and this should be
// no exception.
if (needleLen == 0)
begin++;

// Don't iterate the last needleLen - 1 bytes of the string - they are
// impossible to be splits and would let us memcmp past the end of the
// buffer.
for (const char* iter = begin; iter <= end - needleLen; iter++)
{
// Use of memcmp here instead of strncmp is so that we allow embedded
// nulls to be used in either of the haystack or the needle strings.
// Most Lua string APIs allow embedded nulls, and this should be no
// exception.
if (memcmp(iter, needle, needleLen) == 0)
{
lua_pushinteger(L, ++numMatches);
lua_pushlstring(L, spanStart, iter - spanStart);
lua_settable(L, -3);
// empty separator splits the string into individual characters, so the result size is known up front
lua_createtable(L, int(haystackLen), 0);

spanStart = iter + needleLen;
if (needleLen > 0)
iter += needleLen - 1;
for (const char* iter = haystack; iter < end; iter++)
{
lua_pushlstring(L, iter, 1);
lua_rawseti(L, -2, ++numMatches);
}

return 1;
}
else if (needleLen == 1)
{
// every occurrence of a single character separator is a split, so we can cheaply count them up front
// and allocate the result table at its final size
char sep = needle[0];

if (needleLen > 0)
int count = 1;
for (const char* iter = haystack; (iter = (const char*)memchr(iter, sep, end - iter)) != NULL; iter++)
count++;

lua_createtable(L, count, 0);

for (const char* found; (found = (const char*)memchr(spanStart, sep, end - spanStart)) != NULL; spanStart = found + 1)
{
lua_pushlstring(L, spanStart, found - spanStart);
lua_rawseti(L, -2, ++numMatches);
}
}
else
{
lua_pushinteger(L, ++numMatches);
lua_pushlstring(L, spanStart, end - spanStart);
lua_settable(L, -3);
lua_createtable(L, 0, 0);

if (needleLen <= haystackLen)
{
// Don't iterate the last needleLen - 1 bytes of the string - they are
// impossible to be splits and would let us memcmp past the end of the
// buffer.
const char* last = end - needleLen;

for (const char* iter = haystack; iter <= last;)
{
// the first and the last characters are checked inline to avoid a memcmp call at most positions
if (iter[0] == needle[0] && iter[needleLen - 1] == needle[needleLen - 1] && memcmp(iter, needle, needleLen) == 0)
{
lua_pushlstring(L, spanStart, iter - spanStart);
lua_rawseti(L, -2, ++numMatches);

spanStart = iter + needleLen;
iter = spanStart;
}
else
{
iter++;
}
}
}
}

lua_pushlstring(L, spanStart, end - spanStart);
lua_rawseti(L, -2, ++numMatches);

return 1;
}

Expand Down
61 changes: 61 additions & 0 deletions bench/micro_tests/test_string_split.lua
Original file line number Diff line number Diff line change
@@ -0,0 +1,61 @@
local function prequire(name) local success, result = pcall(require, name); return success and result end
local bench = script and require(script.Parent.bench_support) or prequire("bench_support") or require("../bench_support")

local csv = "alpha,beta,gamma,delta,epsilon,zeta,eta,theta"

local lines = {}
for i = 1, 2000 do
lines[i] = string.rep("x", 40 + i % 40) .. i
end
local text = table.concat(lines, "\n")
local textMulti = table.concat(lines, "\r\n")

local repeated = string.rep("a", 100000)

bench.runCode(function()
local r
for i = 1, 200000 do
r = string.split(csv, ",")
end
assert(#r == 8)
end, "string.split: short csv")

bench.runCode(function()
local r
for i = 1, 100 do
r = string.split(text, "\n")
end
assert(#r == 2000)
end, "string.split: long lines")

bench.runCode(function()
local r
for i = 1, 100 do
r = string.split(textMulti, "\r\n")
end
assert(#r == 2000)
end, "string.split: long lines, 2-char separator")

bench.runCode(function()
local r
for i = 1, 20000 do
r = string.split(csv, "")
end
assert(#r == #csv)
end, "string.split: empty separator")

bench.runCode(function()
local r
for i = 1, 100 do
r = string.split(repeated, "aab")
end
assert(#r == 1)
end, "string.split: separator prefix repeats")

bench.runCode(function()
local r
for i = 1, 100 do
r = string.split(repeated, "aaba")
end
assert(#r == 1)
end, "string.split: separator prefix and suffix repeat")
24 changes: 24 additions & 0 deletions tests/conformance/strings.luau
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,30 @@ do
assert(eq(string.split("abc", "b"), {'a', 'c'}))
assert(eq(string.split("abc", "d"), {'abc'}))
assert(eq(string.split("abc", "c"), {'ab', ''}))
assert(eq(string.split("a,b,c"), {'a', 'b', 'c'}))
assert(eq(string.split("", ""), {}))
assert(eq(string.split("", ","), {''}))
assert(eq(string.split(",", ","), {'', ''}))
assert(eq(string.split(",a,,b,", ","), {'', 'a', '', 'b', ''}))
assert(eq(string.split("abc", "abcd"), {'abc'}))
assert(eq(string.split("abc", "abc"), {'', ''}))
assert(eq(string.split("aaaa", "aa"), {'', '', ''}))
assert(eq(string.split("aaa", "aa"), {'', 'a'}))
assert(eq(string.split("abababa", "aba"), {'', 'b', ''}))
assert(eq(string.split("aabaab", "ab"), {'a', 'a', ''}))
assert(eq(string.split("aXa,aba", "aba"), {'aXa,', ''}))
assert(eq(string.split("aaaa", "aab"), {'aaaa'}))
assert(eq(string.split("a\r\nb\rc\n\r\nd", "\r\n"), {'a', 'b\rc\n', 'd'}))
assert(eq(string.split("a\0b\0\0c", "\0"), {'a', 'b', '', 'c'}))
assert(eq(string.split("a\0\0b\0c", "\0\0"), {'a', 'b\0c'}))
assert(eq(string.split("a\0b", ""), {'a', '\0', 'b'}))

local long = string.rep("x", 1000)
local parts = string.split(table.concat(table.create(100, long), ";"), ";")
assert(#parts == 100)
for i=1,#parts do
assert(parts[i] == long)
end
end

-- validate that variadic string fast calls get correct number of arguments
Expand Down
Loading